- Rename assets/ to docs/, split into guides/ and developer/ - Add get-started.md: installation + 5-step quickstart - Add guides/evaluation.md: 7 eval scripts with CLI args - Add guides/distributed.md: DDP/FSDP, gradient accumulation, NCCL - Add developer/internals.md: loss formulas, RoPE, KV cache math - Add developer/cuda_kernels.md: build system, benchmarks, file layout - Fix storage_format doc in preprocessing.md - Update cross-references in README.md, README-zh-CN.md, Dockerfile
9.1 KiB
Distributed Training
AstrAI supports three parallel modes: single GPU (none), Data Parallel (ddp), and Fully Sharded Data Parallel (fsdp). This guide covers when to use each, how to launch multi-GPU training, and how gradient accumulation works.
Quick Start
Single GPU
python scripts/tools/train.py \
--train_type=sft \
--param_path ./params \
--data_root_path ./dataset \
--parallel_mode=none \
--nprocs=1 \
--batch_per_device=4 \
--grad_accum_steps=8
Multi-GPU DDP (4 GPUs)
export CUDA_VISIBLE_DEVICES=0,1,2,3
export NCCL_P2P_DISABLE=1
export NCCL_NET_GDR_LEVEL=0
python scripts/tools/train.py \
--train_type=sft \
--param_path ./params \
--data_root_path ./dataset \
--parallel_mode=ddp \
--nprocs=4 \
--batch_per_device=4 \
--grad_accum_steps=8
Multi-GPU FSDP (4 GPUs)
export CUDA_VISIBLE_DEVICES=0,1,2,3
export NCCL_P2P_DISABLE=1
export NCCL_NET_GDR_LEVEL=0
python scripts/tools/train.py \
--train_type=sft \
--param_path ./params \
--data_root_path ./dataset \
--parallel_mode=fsdp \
--nprocs=4 \
--batch_per_device=4 \
--grad_accum_steps=8
--parallel_modedefaults tofsdp. You can omit it for FSDP.
Parallel Modes
| Mode | --parallel_mode |
Param Layout | Memory | When to Use |
|---|---|---|---|---|
| Single GPU | none |
Full, replicated | Highest | Small models, DPO/GRPO, debugging |
| DDP | ddp |
Full, replicated | High | Most multi-GPU training |
| FSDP | fsdp |
Sharded (DTensor) | Lowest | Large models that don't fit in single GPU |
NoneExecutor
No wrapping. The model runs as-is on a single device. Gradient accumulation still works via AccumOptimizer/AccumScheduler (they gate step() on the sync counter). Checkpoint saving is a plain state_dict() call.
DDPExecutor
Wraps the model with torch.nn.parallel.DistributedDataParallel. Each rank has a full copy of the model; gradients are all-reduced across ranks. Uses gradient_as_bucket_view=True and broadcast_buffers=False by default (hardcoded in train.py).
During gradient accumulation, non-sync micro-steps use model.no_sync() to skip gradient all-reduce. Only the final micro-step triggers the all-reduce.
FSDPExecutor (FSDP2 / fully_shard)
Uses PyTorch's FSDP2 per-module API (torch.distributed.fsdp.fully_shard). Each model child (e.g., each DecoderBlock) is individually sharded — parameters become DTensors distributed across ranks. No FlatParameter, original parameter names are preserved.
Key differences from DDP:
- Lower memory: parameters are sharded, not replicated.
- Custom grad norm: FSDP gradients are
DTensors, soclip_grad_normcomputes the local norm, then all-reduces to get the global norm. - Collective checkpoint ops:
unshard()andfull_tensor()are collective — all ranks must call them even though only rank-0 saves. The executor handles this viadist.barrier()incheckpoint_context. - Root skipped:
fully_shardis applied to direct children only (not the root model) due to anABC + Generic[T]MRO incompatibility.
Gradient Accumulation
Gradient accumulation lets you simulate a larger effective batch size by accumulating gradients over multiple micro-batches before calling optimizer.step().
Effective batch = nprocs × batch_per_device × grad_accum_steps
Example: 4 GPUs × batch 4 × accum 8 = effective batch 256.
How it works
Three cooperating layers:
GradientState— tracks the micro-step counter. Firessync_gradients=Trueeverygrad_accum_stepsmicro-batches.executor._no_sync(model)— suppresses gradient synchronization on non-sync micro-steps:none:nullcontext(nothing to skip)ddp:model.no_sync()(skips all-reduce)fsdp:set_requires_gradient_sync(False)on eachFSDPModule
AccumOptimizer/AccumScheduler— gatestep()andzero_grad()onsync_gradients, so the optimizer only fires on the last micro-step.
The loss is divided by grad_accum_steps before backward(), so gradients sum to the correct mean.
Process Launching
AstrAI auto-detects the launch method:
| Detection | Strategy | Use Case |
|---|---|---|
torchelastic / torchrun env vars |
TorchrunStrategy |
External orchestrator (torchrun, SLURM, K8s) |
RANK + WORLD_SIZE env vars |
TorchrunStrategy |
External launch |
| Neither | LocalStrategy |
python scripts/tools/train.py (in-process spawn) |
Local (default)
When you run python scripts/tools/train.py --nprocs=4, AstrAI uses torch.multiprocessing.start_processes to spawn 4 child processes. The parent process manages signal forwarding (SIGTERM/SIGINT) and waits for all children to finish.
Torchrun
For multi-node or SLURM environments:
torchrun --nproc_per_node=4 scripts/tools/train.py \
--train_type=sft \
--parallel_mode=ddp \
--param_path ./params \
--data_root_path ./dataset \
--batch_per_device=4
When launched via torchrun, AstrAI reads RANK, WORLD_SIZE, LOCAL_RANK from the environment and uses TorchrunStrategy. The --nprocs flag is ignored (the orchestrator controls process count).
NCCL Environment Variables
For multi-GPU training, you must set these environment variables:
export NCCL_P2P_DISABLE=1
export NCCL_NET_GDR_LEVEL=0
These are required on certain hardware configurations (see AGENTS.md). Without them, NCCL may hang or crash during collective operations. These are set in the training shell scripts (train-seq.sh, train-sft.sh, train-dpo.sh) but not in Python code — you must export them before launching.
Checkpoint Saving
Checkpoints are saved by rank-0 only. The flow:
executor.checkpoint_context(model)— wraps withdist.barrier()before and after (distributed only).executor.unwrap_model(model)— gathers the full state dict:none:model.state_dict()ddp:model.module.state_dict()fsdp:unshard()→full_tensor()→reshard()(collective on all ranks, result kept only on rank-0)
- Non-rank-0 ranks get
None— the save is skipped. - Rank-0 writes
meta.json,config.json,model.safetensors, and optional{key}.pt(optimizer/scheduler state).
FSDP note: Even though only rank-0 saves, all ranks must participate in
unwrap_modelbecauseunshard()andfull_tensor()are collective operations. The barriers incheckpoint_contextkeep all ranks in lockstep.
Total Steps Calculation
The scheduler's total step count accounts for data-parallel sharding:
samples_per_replica = ceil(dataset_len / nprocs)
batches_per_replica = ceil(samples_per_replica / batch_per_device)
total_steps = (batches_per_replica // grad_accum_steps) * n_epoch
This ensures the LR schedule is correctly scaled regardless of the number of GPUs.
Real Examples
Pretraining (seq, DDP, 4 GPUs)
export CUDA_VISIBLE_DEVICES=0,1,2,3
export NCCL_P2P_DISABLE=1
export NCCL_NET_GDR_LEVEL=0
python scripts/tools/train.py \
--train_type=seq \
--param_path ./params \
--data_root_path ./dataset/cached \
--parallel_mode=ddp \
--nprocs=4 \
--n_epoch=1 \
--max_lr=2e-4 \
--schedule_type=wsd \
--warmup_ratio=0.02 \
--window_size=2048 \
--batch_per_device=4 \
--grad_accum_steps=32 \
--ckpt_interval=2000
# Effective batch = 4 × 4 × 32 = 512
SFT (DDP, 4 GPUs)
python scripts/tools/train.py \
--train_type=sft \
--param_path ./AstrAI-V1-base \
--data_root_path ./dataset/cached_sft \
--parallel_mode=ddp \
--nprocs=4 \
--n_epoch=2 \
--max_lr=2e-5 \
--schedule_type=cosine \
--warmup_ratio=0.02 \
--min_rate=0.05 \
--window_size=2048 \
--batch_per_device=4 \
--grad_accum_steps=8
# Effective batch = 4 × 4 × 8 = 128
DPO (Single GPU)
python scripts/tools/train.py \
--train_type=dpo \
--param_path ./checkpoint/epoch_1_step_6000 \
--data_root_path ./alpaca_dpo.jsonl \
--parallel_mode=none \
--nprocs=1 \
--max_lr=5e-6 \
--schedule_type=cosine \
--warmup_ratio=0.1 \
--min_rate=0.1 \
--window_size=1024 \
--batch_per_device=4 \
--grad_accum_steps=8 \
--dpo_beta=0.1 \
--max_grad_norm=50
CLI Parameters
| Parameter | Default | Description |
|---|---|---|
--nprocs |
1 | Number of GPUs / processes |
--parallel_mode |
fsdp |
none, ddp, or fsdp |
--start_method |
spawn |
Multiprocessing start method (spawn, fork, forkserver) |
--backend |
nccl |
Distributed backend (nccl, gloo) |
--master_addr |
localhost |
Master node address |
--master_port |
29500 |
Master node port |
--device_type |
cuda |
Device type |
--tp_sizeis parsed but not yet wired — tensor parallelism is future work.ColumnParallelLinear/RowParallelLinearexist inastrai/parallel/module.pybut are not used by the model.
Full parameter reference: CLI Reference. Training loop and strategies: Training Guide.
Document Update Time: 2026-07-30