Blog #pytorch #machine-learning #ai
PyTorch DDP for Fine-Tuning Gemma 4 on Multiple GPUs
How PyTorch DistributedDataParallel works and when it is the right tool for LoRA fine-tuning Gemma 4, with a complete torchrun script, memory math for the 12B model, and the failure modes that hang multi-GPU jobs.
- Published
- Reading
- 29 min
- Author
- Sudhanva Narayana
A while back I built a small Megatron Parallelism Visualizer to make the three classic ways of splitting a training job across GPUs concrete: tensor parallelism, pipeline parallelism, and data parallelism. Tensor and pipeline parallelism get most of the attention because they’re what makes a 70B-parameter model trainable at all. Most fine-tuning jobs I see don’t need either of them. They need the plain one.
This post covers the plain one, PyTorch DistributedDataParallel (DDP), with a concrete target: LoRA fine-tuning Gemma 4 on a single multi-GPU machine. It explains how DDP actually moves gradients, gives a complete torchrun script you can adapt, works through the memory budget for the Gemma 4 12B model, and lists the failure modes that turn a multi-GPU job into a silent hang.
A note on what is measured and what isn’t. I haven’t run a production Gemma 4 fine-tune for this post, and there are no loss curves or speedup charts here. The model facts come from Google’s model card and the published checkpoint files. The script was smoke-tested end to end on CPU with two processes and a tiny, randomly initialized model built from the real Gemma 4 12B config and tokenizer (details in What I verified). Every memory and bandwidth figure below is arithmetic, and each one is labeled as an estimate.
Table of contents
- When DDP is the right first tool
- Gemma 4 facts that matter for training
- How DDP works
- Batch size and learning rate
- Memory budget for Gemma 4 12B with LoRA
- The training script
- Launching with torchrun
- Running it on Modal
- Failure modes that hang or corrupt a DDP job
- NCCL and debugging environment variables
- Measuring throughput and scaling
- What I verified and what I didn’t
When DDP is the right first tool
The visualizer’s README puts the three strategies in one line each:
- Tensor parallelism (TP) splits individual weight matrices across GPUs. Each GPU computes a partial matmul, then they communicate to reassemble the result.
- Pipeline parallelism (PP) assigns different layers to different GPUs and streams micro-batches through them, paying for it in pipeline “bubble” time.
- Data parallelism (DP) gives every GPU a full copy of the model and a different slice of the batch, then averages gradients with an all-reduce.
Megatron-Core combines all three (“3D parallelism”) because at frontier scale no single strategy is enough. Fine-tuning one model on one eight-GPU node is a different problem. The question that decides it is simple: does one replica of the training state fit on one GPU?
flowchart TD
Fit{"One replica fits<br/>on one GPU?"}
Fit -->|Yes| DDP["DDP<br/>copy model, split data"]
Fit -->|No| Shard{"Only optimizer state<br/>and grads too big?"}
Shard -->|Yes| FSDP["FSDP / ZeRO<br/>shard the state"]
Shard -->|No| TPPP["Tensor / pipeline<br/>parallelism + DP"]
DDP -.-> Shrink["Almost fits? LoRA, QLoRA,<br/>checkpointing, shorter seqs"]
With LoRA, the answer for most Gemma 4 sizes is yes. The frozen base weights sit in bf16, the trainable adapter is tens of millions of parameters, and gradient checkpointing keeps activations bounded. Once a replica fits, DDP is the simplest correct way to use more GPUs: nothing inside the model changes, the only collective is a gradient all-reduce, and every rank runs the same code.
When a replica doesn’t fit, you have options, and they solve different problems:
| Strategy | What it splits | Reach for it when | Cost |
|---|---|---|---|
| DDP | The data | A full replica fits on one GPU | One gradient all-reduce per optimizer step |
| FSDP / ZeRO | Parameters, gradients, optimizer state | Full fine-tuning, where optimizer state dwarfs the weights | All-gathers in forward and backward, reduce-scatter of grads |
| Tensor parallelism | Individual weight matrices | Single layers are too large, or you need lower per-step latency | Collectives inside every layer, needs fast intra-node links |
| Pipeline parallelism | Layers | The model’s depth doesn’t fit, especially across nodes | Pipeline bubbles, micro-batch scheduling complexity |
For full fine-tuning Gemma 4 12B the table points at FSDP, not DDP. The memory section shows why. For LoRA on the same model, DDP is the right tool.
Gemma 4 facts that matter for training
Google released Gemma 4 in four sizes on March 31, 2026, and added a 12B model on June 3, 2026. The details that matter for a training plan, from the Gemma 4 model card and the Gemma 4 model overview:
| Model | Parameters | Layers | Context | Input modalities | BF16 inference memory (Google) |
|---|---|---|---|---|---|
| E2B | 2.3B effective, 5.1B with embeddings | 35 | 128K | Text, image, audio | 11.4 GB |
| E4B | 4.5B effective, 8B with embeddings | 42 | 128K | Text, image, audio | 17.9 GB |
| 12B | 11.95B | 48 | 256K | Text, image, audio | 26.7 GB |
| 26B A4B | 25.2B total, 3.8B active (MoE) | 30 | 256K | Text, image | 57.7 GB |
| 31B | 30.7B | 60 | 256K | Text, image | 69.9 GB |
A few more points from the primary sources:
- License. The weights are released under Apache 2.0, per the model card, the Gemma 4 launch post, and the Hugging Face repository for gemma-4-12B-it.
- Model IDs. Base checkpoints are
google/gemma-4-E2B,google/gemma-4-E4B,google/gemma-4-12B,google/gemma-4-26B-A4B, andgoogle/gemma-4-31B. Instruction-tuned variants append-it. - Tooling. The launch post lists day-one support in Hugging Face Transformers and TRL, among many others. Google’s own text fine-tuning guide for Gemma uses Transformers, PEFT, and TRL, installs
transformers>=5.10.1andpeft>=0.19.0, and trains a LoRA adapter (QLoRA in that guide) on a single 16 GB T4. - Architecture. All sizes use a 262K vocabulary and interleave sliding-window attention with global attention layers. The 12B model is encoder-free: image and audio inputs go straight into the decoder instead of through separate encoders.
I also read the published checkpoint headers and configs for google/gemma-4-12B-it, because a few details there change the training plan:
- 11.96B parameters, all stored in BF16, and 11.91B of them live under
model.language_model. The input embedding alone is262144 × 3840, about 1.0B parameters, tied to the output head. - Eight of the 48 layers are global attention (every sixth layer). Those layers set
attention_k_eq_v, so they have ak_projbut nov_proj. A LoRA config that listsv_projstill works, it just has nothing to attach to in those eight layers. - LoRA at rank 16 on every attention and MLP projection of the text decoder is 65.6M trainable parameters, or about 0.55% of the model. I counted that from the tensor shapes.
For this post I use google/gemma-4-12B-it. It is big enough that the memory budget matters, and small enough that a bf16 LoRA replica fits on a single 48 GB or 80 GB GPU, which is exactly when DDP is the right tool.
How DDP works
The visualizer’s data-parallel simulation, simulate_data_parallel in simulator.py, is the mental model: every GPU runs forward and backward on its own slice of the batch, across all gradient-accumulation steps, and then one all-reduce averages the gradients so every replica applies the same update. The dataParallel client simulator in simulator.js animates the same sequence.
Real DDP follows that shape, with one important refinement the simulation leaves out. Here is the life of a DDP job, following the PyTorch DDP design notes and the DistributedDataParallel API reference.
1. One process per GPU
torchrun starts one Python process per GPU and sets RANK, LOCAL_RANK, and WORLD_SIZE in each process’s environment. Each process pins itself to its GPU and joins a process group. With the NCCL backend, the collectives run on the GPUs over NVLink or PCIe within a node, and over the network across nodes.
2. Replicas start identical
When you wrap a model in DistributedDataParallel, the constructor broadcasts parameters and buffers from rank 0 to every other rank. All replicas begin from the same weights, including freshly initialized LoRA matrices. DDP also registers an autograd hook on every parameter that requires a gradient. Frozen parameters get no hook and never cross the wire. That’s why LoRA and DDP combine so cheaply.
3. The sampler shards the data
DistributedSampler gives each rank a disjoint slice of the dataset indices. It shuffles with a seed plus the epoch number, so call sampler.set_epoch(epoch) every epoch. Without it, every epoch replays the same order. By default the sampler pads the index list so every rank gets the same number of samples. With drop_last=True it trims instead. Either way every rank sees the same number of batches, and that matters more than it sounds (see failure modes).
4. Backward overlaps with communication
This is the part the simulation simplifies away. The simulation runs forward, then backward, then all-reduce as separate phases. Real DDP groups parameters into buckets (the default bucket_cap_mb is 25 MB), roughly in reverse order of the model’s parameters, since that’s about the order in which gradients become ready during backward. When every gradient in a bucket has been computed, DDP launches an asynchronous all-reduce for that bucket while backward keeps computing the earlier layers. By the time autograd reaches the first layer, most of the communication has already finished.
sequenceDiagram
participant G0 as GPU 0
participant G1 as GPU 1
participant NCCL as NCCL
Note over G0,G1: micro-batches 1..N-1 in no_sync()
G0->>G0: forward (last micro-batch)
G1->>G1: forward (last micro-batch)
G0->>G0: backward: last layers
G1->>G1: backward: last layers
G0->>NCCL: bucket k ready
G1->>NCCL: bucket k ready
Note over NCCL: all-reduce bucket k
G0->>G0: backward: earlier layers
G1->>G1: backward: earlier layers
G0->>NCCL: bucket 1 ready
G1->>NCCL: bucket 1 ready
NCCL-->>G0: averaged gradients
NCCL-->>G1: averaged gradients
Note over G0,G1: clip + optimizer.step()
Gradients are averaged across ranks, not summed. The optimizer step then runs independently on every rank, from identical gradients and identical starting weights, so the replicas stay identical without ever broadcasting weights again.
5. no_sync() for gradient accumulation
With gradient accumulation you want the all-reduce once per optimizer step, not once per micro-batch. model.no_sync() is a context manager that skips the reduction. Gradients accumulate locally in .grad, and the first forward and backward pass outside the context synchronizes everything accumulated so far. The API docs carry a warning worth repeating: the forward pass has to run inside the context too, or gradients still get synchronized. This is exactly what the visualizer draws, one all-reduce after all the accumulation steps. Its step comment even reads # All-reduce after accumulation.
6. What DDP doesn’t do
DDP only averages gradients. It doesn’t average your loss for logging, doesn’t coordinate checkpoint writes, doesn’t make evaluation distributed, and doesn’t protect you from ranks that take different code paths. All of that is on you, which is most of what the script below handles.
Batch size and learning rate
The effective (global) batch size is:
global_batch = micro_batch_per_gpu × world_size × accumulation_steps
The simulator states the same identity in slightly different terms. Its batch_size argument is the per-step batch across all GPUs, so micro_batch_size = batch_size // num_gpus and effective_batch_size = batch_size × gradient_accumulation_steps. The data-parallel tests pin it down: 128 samples on 2 GPUs with 4 accumulation steps gives an effective batch of 512. One nit in my own code: with a batch that doesn’t divide evenly, the integer division drops the remainder but effective_batch_size still reports the requested size. In a real training script, always derive the global batch from the micro-batch, which is what the script below prints.
Two consequences for fine-tuning:
- Adding GPUs changes the optimization problem. Going from 1 GPU to 8 with the same per-GPU micro-batch and accumulation multiplies the global batch by 8 and cuts the number of optimizer steps per epoch by 8. If you tuned a learning rate on one GPU, that tuning doesn’t carry over automatically. The DDP docs say the same thing: because gradients are averaged, a DDP model behaves like a single-GPU model trained on the global batch.
- Pick the global batch first, then solve for accumulation. If the recipe you trust used a global batch of 64, keep 64: on 8 GPUs with a micro-batch of 1, that’s 8 accumulation steps. Then the only thing DDP changes is wall-clock time.
The linear scaling rule (scale the learning rate with the batch, plus warmup) came from large-batch SGD on ImageNet. For Adam-style optimizers on LoRA fine-tunes it’s a starting heuristic, not a law. Treat the learning rate as something you sweep at the global batch you actually run.
One subtlety specific to language models: Hugging Face’s causal-LM loss averages over the tokens in each micro-batch, and DDP then averages equally across ranks. Micro-batches with very different token counts therefore get equal weight regardless of how many tokens they hold. For chat fine-tuning with variable-length samples this is usually fine. If it matters for your data, sum the per-token loss, all-reduce the token count, and divide by the global token count.
Memory budget for Gemma 4 12B with LoRA
Everything in this section is an estimate, meant to decide which GPUs to request before you pay for them. Measure on your hardware with torch.cuda.max_memory_allocated() after the first few steps. GB means 10⁹ bytes.
Assumptions: google/gemma-4-12B-it, bf16 base weights, LoRA rank 16 on all text-decoder projections (65.6M trainable parameters), adapter weights in fp32 (PEFT’s default when the base model is bf16), AdamW, gradient checkpointing on, micro-batch 1, sequence length 2,048.
| Component | Arithmetic | Estimate |
|---|---|---|
| Frozen base weights (bf16) | 11.96B × 2 bytes | 23.9 GB |
| LoRA weights (fp32) | 65.6M × 4 bytes | 0.26 GB |
| LoRA gradients (fp32) | 65.6M × 4 bytes | 0.26 GB |
| AdamW moments (fp32, two per param) | 65.6M × 8 bytes | 0.52 GB |
| DDP gradient buckets | 0 with gradient_as_bucket_view=True, else +0.26 GB |
0–0.26 GB |
| Checkpointed layer inputs | 2,048 tokens × 3,840 hidden × 2 bytes × 48 layers | 0.75 GB |
| One layer’s recomputed activations | MLP intermediates at 15,360 wide dominate | ~1 GB |
| Logits and loss | 2,048 × 262,144 = 537M values; bf16 logits, softcap temporaries, fp32 upcast for cross-entropy, and its gradient | ~6–9 GB |
| CUDA context, NCCL buffers, allocator slack | Varies by driver and fragmentation | ~2–4 GB |
| Total per GPU | ~35–40 GB |
Three things stand out.
The trainable state is small. Adapter weights, gradients, and optimizer moments together come to about 1 GB. That’s what makes DDP viable: every GPU holds a full copy of them and it barely matters.
The vocabulary is the activation hog. With 262K entries, the logits tensor at 2,048 tokens is 537M values, and cross-entropy wants it in fp32. Gemma 4 also applies a final logit softcap (tanh scaling), which creates extra temporaries. Doubling the sequence length to 4,096 roughly doubles this line, to around 12–18 GB. When you run out of memory at long sequence lengths, this is usually where to look first, not the decoder layers.
The 16-bytes-per-parameter rule rules out DDP for full fine-tuning. Mixed-precision Adam holds about 16 bytes per trainable parameter: bf16 weights and gradients, plus fp32 master weights and two fp32 moments, as the ZeRO paper lays out. For 11.96B parameters that’s about 191 GB before any activations, which fits on no single GPU. Full fine-tuning the 12B model means sharding that state with FSDP or ZeRO, not replicating it with DDP.
The same lens for the other sizes (weights only, bf16, my arithmetic):
| Model | bf16 weights | DDP + LoRA outlook (estimate) |
|---|---|---|
| E4B | 8.0B × 2 ≈ 16 GB | Fits 40–80 GB GPUs comfortably. On 24 GB GPUs, shorten sequences since the logits line is the same size as for 12B. |
| 12B | 11.96B × 2 ≈ 24 GB | Fits 48 GB GPUs at 2K tokens, 80 GB with headroom for longer sequences or micro-batch 2. |
| 26B A4B | 25.2B × 2 ≈ 50 GB | Memory scales with total parameters, not active ones. Tight on 80 GB. Consider QLoRA or FSDP. |
| 31B | 30.7B × 2 ≈ 61 GB | Too tight in bf16 on 80 GB once activations are added. QLoRA (4-bit base) per replica, or FSDP. |
For the MoE model, the 3.8B active parameters set the compute per token, but every expert still has to be resident on every replica.
Communication per optimizer step
DDP moves gradients, so LoRA shrinks the network cost by the same factor it shrinks the trainable state. A ring all-reduce sends about 2 × (N − 1) / N times the buffer size per GPU, which is the bus-bandwidth factor nccl-tests uses in its performance notes. For N = 8, that factor is 1.75.
| What you train | Gradient buffer | Sent per GPU per step (estimate) |
|---|---|---|
| LoRA r=16, text decoder (this post) | 65.6M × 4 bytes ≈ 0.26 GB | ≈ 0.46 GB |
LoRA plus trainable embed_tokens / lm_head (tied, ~1.0B params) |
+2.0 GB (bf16) to +4.0 GB (fp32) | ≈ 4–7.5 GB |
| Full fine-tune, bf16 gradients | 11.96B × 2 bytes ≈ 23.9 GB | ≈ 42 GB |
The middle row is worth knowing about. Google’s fine-tuning guide sets modules_to_save=["lm_head", "embed_tokens"] in its LoRA config. On a single GPU that’s a memory question. Under DDP it also makes the all-reduce about ten times larger. If you don’t need the embeddings to move, leave them frozen.
To turn bytes into time, divide by the bus bandwidth that nccl-tests measures on your own machine. As an illustration only: at an assumed 100 GB/s of bus bandwidth, the LoRA all-reduce would take on the order of 5 ms, once per optimizer step, and most of it hides behind backward anyway.
The training script
This is a complete, minimal DDP training script. It uses plain PyTorch for the distributed parts, Transformers to load Gemma 4, and PEFT for LoRA. I left out the Hugging Face Trainer and TRL on purpose, so every distributed concept is visible. In production I’d usually reach for TRL’s SFTTrainer launched under torchrun, which does the same things underneath.
The dataset format is the one Google’s guide uses, one JSON object per line:
{
"messages": [
{ "role": "user", "content": "..." },
{ "role": "assistant", "content": "..." }
]
}
# train_ddp.py: LoRA fine-tuning of Gemma 4 with plain PyTorch DDP.
# Launch: torchrun --standalone --nproc_per_node=8 train_ddp.py
import contextlib
import json
import os
import time
import torch
import torch.distributed as dist
from peft import LoraConfig, get_peft_model
from torch.nn.parallel import DistributedDataParallel as DDP
from torch.utils.data import DataLoader, Dataset
from torch.utils.data.distributed import DistributedSampler
from transformers import (
AutoModelForCausalLM,
AutoTokenizer,
get_cosine_schedule_with_warmup,
)
MODEL_ID = os.environ.get("MODEL_ID", "google/gemma-4-12B-it")
DATA_PATH = os.environ.get("DATA_PATH", "train.jsonl")
OUTPUT_DIR = os.environ.get("OUTPUT_DIR", "checkpoints")
MAX_LEN = 2048 # tokens per sample after truncation
MICRO_BATCH = 1 # samples per GPU per forward pass
ACCUM_STEPS = 8 # micro-batches per optimizer step
EPOCHS = 1
LR = 1e-4
WARMUP_STEPS = 20
MAX_GRAD_NORM = 1.0
LOG_EVERY = 10
SEED = 42
# Only the text decoder's projections. Keeps LoRA off the vision/audio paths,
# which a text-only batch never touches.
LORA_TARGETS = (
r".*language_model.*\.(q_proj|k_proj|v_proj|o_proj|gate_proj|up_proj|down_proj)"
)
class ChatJsonl(Dataset):
"""One JSON object per line: {"messages": [{"role": ..., "content": ...}]}."""
def __init__(self, path, tokenizer, max_len):
with open(path) as f:
self.rows = [json.loads(line) for line in f if line.strip()]
self.tokenizer = tokenizer
self.max_len = max_len
def __len__(self):
return len(self.rows)
def __getitem__(self, idx):
text = self.tokenizer.apply_chat_template(
self.rows[idx]["messages"], tokenize=False
)
# The chat template already emits <bos>, so don't add it twice.
ids = self.tokenizer(
text, add_special_tokens=False, truncation=True, max_length=self.max_len
)["input_ids"]
return torch.tensor(ids, dtype=torch.long)
class PadCollate:
"""Right-pads a batch and masks padding out of the loss.
A top-level class, not a closure, so DataLoader workers can pickle it.
"""
def __init__(self, pad_id):
self.pad_id = pad_id
def __call__(self, batch):
longest = max(len(x) for x in batch)
input_ids = torch.full((len(batch), longest), self.pad_id, dtype=torch.long)
attention_mask = torch.zeros((len(batch), longest), dtype=torch.long)
for i, ids in enumerate(batch):
input_ids[i, : len(ids)] = ids
attention_mask[i, : len(ids)] = 1
labels = input_ids.masked_fill(attention_mask == 0, -100)
return {
"input_ids": input_ids,
"attention_mask": attention_mask,
"labels": labels,
}
def main():
# torchrun sets RANK, LOCAL_RANK, WORLD_SIZE, MASTER_ADDR and MASTER_PORT.
local_rank = int(os.environ["LOCAL_RANK"])
torch.cuda.set_device(local_rank)
device = torch.device("cuda", local_rank)
dist.init_process_group(backend="nccl", device_id=device)
rank = dist.get_rank()
world_size = dist.get_world_size()
is_main = rank == 0
torch.manual_seed(SEED) # same LoRA init on every rank
tokenizer = AutoTokenizer.from_pretrained(MODEL_ID)
model = AutoModelForCausalLM.from_pretrained(
MODEL_ID, dtype=torch.bfloat16, device_map={"": local_rank}
)
model.gradient_checkpointing_enable(
gradient_checkpointing_kwargs={"use_reentrant": False}
)
model = get_peft_model(
model,
LoraConfig(
r=16,
lora_alpha=32,
lora_dropout=0.05,
target_modules=LORA_TARGETS,
task_type="CAUSAL_LM",
),
)
if is_main:
model.print_trainable_parameters()
model = DDP(model, device_ids=[local_rank], gradient_as_bucket_view=True)
torch.manual_seed(SEED + rank) # but different dropout masks per rank
dataset = ChatJsonl(DATA_PATH, tokenizer, MAX_LEN)
sampler = DistributedSampler(dataset, shuffle=True, seed=SEED, drop_last=True)
loader = DataLoader(
dataset,
batch_size=MICRO_BATCH,
sampler=sampler,
collate_fn=PadCollate(tokenizer.pad_token_id),
drop_last=True,
num_workers=2,
pin_memory=True,
)
trainable = [p for p in model.parameters() if p.requires_grad]
optimizer = torch.optim.AdamW(trainable, lr=LR, weight_decay=0.0, fused=True)
steps_per_epoch = len(loader) // ACCUM_STEPS # identical on every rank
scheduler = get_cosine_schedule_with_warmup(
optimizer, WARMUP_STEPS, steps_per_epoch * EPOCHS
)
if is_main:
global_batch = MICRO_BATCH * world_size * ACCUM_STEPS
print(f"world={world_size} global_batch={global_batch} steps={steps_per_epoch}")
for epoch in range(EPOCHS):
sampler.set_epoch(epoch) # new shuffle each epoch, same on all ranks
model.train()
batches = iter(loader)
for step in range(steps_per_epoch):
t0 = time.perf_counter()
loss_sum = torch.zeros((), device=device)
tokens = torch.zeros((), device=device)
for micro in range(ACCUM_STEPS):
batch = {
k: v.to(device, non_blocking=True) for k, v in next(batches).items()
}
last = micro == ACCUM_STEPS - 1
# Skip the gradient all-reduce on all but the last micro-batch.
sync = contextlib.nullcontext() if last else model.no_sync()
with sync:
with torch.autocast("cuda", dtype=torch.bfloat16):
out = model(**batch, use_cache=False)
loss = out.loss / ACCUM_STEPS
loss.backward()
loss_sum += loss.detach()
tokens += batch["attention_mask"].sum()
grad_norm = torch.nn.utils.clip_grad_norm_(trainable, MAX_GRAD_NORM)
optimizer.step()
scheduler.step()
optimizer.zero_grad(set_to_none=True)
if step % LOG_EVERY == 0:
# Logging only: training already synced gradients, not losses.
stats = torch.stack([loss_sum, tokens])
dist.all_reduce(stats, op=dist.ReduceOp.SUM)
loss_avg = stats[0].item() / world_size
total_tokens = stats[1].item()
dt = time.perf_counter() - t0 # .item() above waited for the GPU
if is_main:
print(
f"epoch={epoch} step={step} loss={loss_avg:.4f} "
f"grad_norm={grad_norm.item():.3f} "
f"lr={scheduler.get_last_lr()[0]:.2e} "
f"tok/s={total_tokens / dt:,.0f}"
)
if is_main:
ckpt = os.path.join(OUTPUT_DIR, f"epoch-{epoch}")
model.module.save_pretrained(ckpt) # adapter weights only
torch.save(
{
"optimizer": optimizer.state_dict(),
"scheduler": scheduler.state_dict(),
"epoch": epoch,
},
os.path.join(ckpt, "trainer_state.pt"),
)
dist.barrier() # nobody races ahead while rank 0 writes
dist.destroy_process_group()
if __name__ == "__main__":
main()
Why each piece is there
Process group and device. torch.cuda.set_device(local_rank) comes before anything touches CUDA, so each process allocates only on its own GPU. Passing device_id to init_process_group binds the process group to that device up front. LOCAL_RANK picks the GPU on this machine. RANK (via dist.get_rank()) is the global identity used for “only rank 0 does this.”
Loading the model. device_map={"": local_rank} loads the whole model straight onto this rank’s GPU. Every rank loads its own copy, which is the point of DDP. Download the checkpoint once before launching (hf download google/gemma-4-12B-it), or eight processes will race to fill the same cache. On current Transformers, AutoModelForCausalLM resolves the 12B checkpoint to its full multimodal class. The model card uses AutoModelForMultimodalLM, and either is fine for text-only training.
LoRA targets. LORA_TARGETS is a regex (PEFT treats a string target_modules as a regex match) that only matches projections under language_model. For a text-only dataset, adapters on image or audio paths would never receive gradients, and DDP raises an error when a parameter that requires a gradient doesn’t get one. Scoping the adapters is cleaner than switching on find_unused_parameters=True and paying for an extra graph traversal on every step.
Gradient checkpointing. Non-reentrant checkpointing (use_reentrant=False) is the variant that composes cleanly with DDP and with frozen base weights. The reentrant variant has well-known sharp edges with DDP. Checkpointing is what keeps the activation lines in the memory table small.
Seeds. Seeding identically before get_peft_model gives every rank the same adapter initialization (DDP’s rank-0 broadcast would also enforce that). Reseeding with SEED + rank after the wrap gives each rank different LoRA dropout masks, which it should have.
gradient_as_bucket_view=True. Gradients become views into DDP’s communication buckets instead of separate tensors that get copied in and out. That saves one gradient-sized buffer and a copy per step. It’s small for LoRA and substantial for anything bigger.
Sampler and loader. DistributedSampler(drop_last=True) plus DataLoader(drop_last=True) means every rank gets the same number of full micro-batches. steps_per_epoch is computed from len(loader), which is identical on every rank, and the loop runs exactly that many optimizer steps. A trailing partial accumulation window is dropped instead of being handled differently on different ranks.
Accumulation with no_sync(). The first ACCUM_STEPS − 1 micro-batches run forward and backward inside model.no_sync(), then the last one runs outside it, which triggers the bucketed all-reduce during its backward. Dividing the loss by ACCUM_STEPS makes the accumulated gradient a mean rather than a sum.
bf16 autocast. The base weights are already bf16. The LoRA weights are fp32, because PEFT upcasts bf16 adapters by default for stable training. Autocast runs the matmuls in bf16 while the optimizer updates fp32 adapter weights. No GradScaler is needed: that’s for fp16, whose narrow exponent range underflows. bf16 has the same exponent range as fp32.
Clipping. Clipping runs after the all-reduce, on identical averaged gradients, so every rank computes the same norm and applies the same scale.
Logging. The logged loss is an explicit all_reduce of the local losses, run on every rank. LOG_EVERY is rank-independent. The .item() calls force a sync with the GPU, so the step time is honest instead of measuring only how long it took to queue kernels.
Checkpointing. Replicas are identical, so only rank 0 writes. model.module unwraps DDP, and save_pretrained on the PEFT model writes only the adapter (about 262 MB for this config in fp32). Optimizer state is identical across ranks too, so one copy is enough to resume. The barrier() keeps the other ranks from racing ahead, or exiting, while rank 0 writes.
Teardown. destroy_process_group() shuts down NCCL cleanly, so the job exits instead of hanging on shutdown or printing warnings.
What I’d add before a real run
- Loss masking to assistant tokens only. The script trains on the whole rendered conversation, prompts included. Most chat fine-tunes mask everything except assistant turns by setting those label positions to
-100. - Resume. Load the adapter with PEFT and the optimizer and scheduler state from
trainer_state.pt, and fast-forward the sampler to the saved epoch. - Evaluation. Either run it on every rank over a
DistributedSamplerand all-reduce the metrics, or run it on rank 0 with the unwrappedmodel.moduleundertorch.no_grad()while the others wait at a barrier. Never call the DDP-wrapped model’s forward on one rank only. - Token-weighted loss if your sample lengths vary wildly (see batch size and learning rate).
Launching with torchrun
On one node with eight GPUs:
pip install torch "transformers>=5.10.1" "peft>=0.19.0" accelerate
hf download google/gemma-4-12B-it
torchrun --standalone --nproc_per_node=8 train_ddp.py
--standalone sets up a local rendezvous, so you don’t have to pick a master address or port. Across two nodes, run the same command on each node with a shared rendezvous endpoint:
torchrun \
--nnodes=2 \
--nproc_per_node=8 \
--rdzv_backend=c10d \
--rdzv_endpoint=node0.example.internal:29400 \
--rdzv_id=gemma4-lora-001 \
train_ddp.py
The script doesn’t change. RANK and WORLD_SIZE now span 16 processes, LOCAL_RANK still runs from 0 to 7 on each node, and rank 0 (the only writer) lives on one node, so point OUTPUT_DIR at storage that rank’s node can write to. The torchrun documentation covers elastic options like --max_restarts.
Across nodes the all-reduce leaves NVLink and goes over the network. For LoRA’s roughly quarter-gigabyte buffer, that’s usually tolerable. For full fine-tuning it’s the reason people buy InfiniBand.
Running it on Modal
The visualizer’s GPU benchmarks run on Modal, where the 2×T4 tensor-parallel endpoint is just @app.function(gpu="T4:2"). The same pattern works for this script. Modal’s GPU guide documents the "H100:8" syntax for eight GPUs in one container and suggests running multi-process training as a subprocess, which is exactly what torchrun wants:
# modal_train.py
import os
import subprocess
import modal
app = modal.App("gemma4-ddp")
volume = modal.Volume.from_name("gemma4-ddp", create_if_missing=True)
image = (
modal.Image.debian_slim(python_version="3.12")
.pip_install("torch", "transformers>=5.10.1", "peft>=0.19.0", "accelerate")
.env({"HF_HOME": "/vol/hf"})
.add_local_file("train_ddp.py", "/root/train_ddp.py")
.add_local_file("train.jsonl", "/root/train.jsonl")
)
@app.function(image=image, gpu="H100:8", volumes={"/vol": volume}, timeout=6 * 60 * 60)
def train():
env = {
**os.environ,
"DATA_PATH": "/root/train.jsonl",
"OUTPUT_DIR": "/vol/checkpoints",
}
subprocess.run(
["torchrun", "--standalone", "--nproc_per_node=8", "/root/train_ddp.py"],
check=True,
env=env,
)
volume.commit()
@app.local_entrypoint()
def main():
train.remote()
Run it with modal run modal_train.py. The volume holds both the Hugging Face cache, so later runs skip the download, and the adapter checkpoints.
Failure modes that hang or corrupt a DDP job
Most DDP bugs don’t crash. They hang: one rank waits in a collective that another rank never enters, until the NCCL watchdog times out. The rule behind every item below is the same: every rank has to issue the same collectives, in the same order, with the same shapes.
Unequal numbers of batches
If rank 3 gets one more batch than rank 5, rank 3 enters an all-reduce that rank 5 never joins. Causes include custom samplers that don’t pad, an IterableDataset sharded by file where files have different lengths, and filtering samples inside the training loop on some ranks.
Fix: Use DistributedSampler with drop_last=True, compute the step count from len(loader) as the script does, and never continue past a batch on one rank only. If uneven inputs are inherent, PyTorch’s Join context manager for uneven inputs lets ranks that finish early shadow the collectives of the ones still running.
Parameters that don’t receive gradients
Symptom: a runtime error on the second iteration saying DDP expected to finish reduction in the prior iteration, listing parameters that didn’t receive gradients.
Cause: A parameter with requires_grad=True wasn’t used in the forward pass. For Gemma 4 the usual culprit is target_modules="all-linear" on a multimodal checkpoint, which attaches LoRA to vision and audio projections that a text-only batch never runs. I reproduced this with the tiny test model built from the 12B config: with all-linear, the first step succeeds and the second fails with Expected to have finished reduction in the prior iteration before starting a new one, naming six parameter indices that never received a gradient.
Fix: Scope the adapters to the modules you actually exercise, like the regex in the script, or freeze the unused modules. find_unused_parameters=True also fixes it, at the price of a full graph traversal after every forward pass. Use it when unused parameters are genuinely data-dependent. That’s plausible with a mixture-of-experts model if your adapters sit on expert weights that some micro-batches never route to.
Rank-divergent code paths
Anything like if rank == 0: dist.all_reduce(...), running evaluation through the DDP wrapper on rank 0 only, or a break taken on one rank because its loss went NaN. All of these desynchronize the collectives.
Fix: Keep collectives unconditional. Make decisions like early stopping on rank 0, broadcast them as a tensor, and have every rank act on the broadcast value.
Changing the model after wrapping
DDP registers its hooks at construction time. Adding adapters, unfreezing layers, or swapping modules after DDP(...) means the hooks no longer match the parameters, and the API docs warn against it. Apply PEFT, gradient checkpointing, and any freezing before wrapping, which is the order the script uses.
Doubled or missing synchronization with accumulation
Running forward outside no_sync() and only backward inside it silently synchronizes every micro-batch. Nothing breaks, it’s just slower. The opposite mistake, running the last micro-batch inside no_sync() too, skips the all-reduce entirely, and the replicas quietly drift apart. A cheap check is to compare a checksum of the trainable parameters across ranks after a few steps. That’s what I did in the smoke test.
DataLoader workers and pickling
With the spawn or forkserver multiprocessing start methods (spawn is the default on macOS, and the DDP docs recommend forkserver or spawn in some NCCL setups), everything a worker receives must be picklable. My first draft used a closure as collate_fn, and it failed immediately under spawn with Can't get local object 'make_collate.<locals>.collate'. That’s why the script uses a top-level PadCollate class.
Every rank writing checkpoints
Eight processes writing the same file produces corrupt or interleaved output on shared filesystems. Write on rank 0, then barrier().
NCCL and debugging environment variables
These are the environment variables I set before digging into a distributed problem. Details are in NVIDIA’s NCCL environment variable reference and the PyTorch distributed documentation.
| Variable | What it does |
|---|---|
NCCL_DEBUG=INFO |
Logs NCCL’s init: which transports (NVLink, PCIe, InfiniBand, sockets) it picked and the ring/tree topology. |
NCCL_DEBUG_SUBSYS=INIT,NET |
Narrows the log to initialization and networking. |
NCCL_SOCKET_IFNAME=eth0 |
Pins the network interface. Multi-node jobs often hang because NCCL picked a Docker bridge or the wrong NIC. |
NCCL_IB_DISABLE=1 |
Disables InfiniBand, to test whether a hang is IB-specific. A diagnostic step, not a fix. |
TORCH_DISTRIBUTED_DEBUG=DETAIL |
Makes PyTorch check collective consistency and report which parameters didn’t receive gradients. |
TORCH_NCCL_TRACE_BUFFER_SIZE=2000 |
Enables PyTorch’s NCCL flight recorder, which records recent collectives per rank so you can see who stopped where. |
Two more habits help. Run nvidia-smi topo -m once per machine type to see whether GPUs talk over NVLink or PCIe, since that changes what throughput to expect. And pass a shorter timeout to init_process_group while debugging, so a hang turns into a stack trace in minutes instead of the default wait.
Measuring throughput and scaling
The number to watch is tokens per second, across all ranks and per GPU. The script logs non-padding tokens per second for each logged step. To measure scaling:
- Run the same config on 1, 2, 4, and 8 GPUs, keeping the micro-batch per GPU fixed.
- Discard the first steps (CUDA context setup, allocator warmup, and NCCL communicator setup all land there) and average over a steady window.
- Compute scaling efficiency as
tokens_per_s(N) / (N × tokens_per_s(1)).
That last formula is the same one the visualizer’s 2×T4 benchmark reports. api_multi_gpu_benchmark in app.py computes parallel_efficiency = speedup / ideal_speedup × 100 for a real column-parallel matmul split across two T4s, and its breakdown field separates compute time from GPU-to-GPU transfer time. The README doesn’t publish a fixed result, because the endpoint measures live on each call. The lesson it’s built to show carries over directly: splitting work across GPUs speeds it up, and communication eats part of that speedup.
What to expect from DDP with LoRA, qualitatively:
- The all-reduce is small and mostly hidden. At about 0.26 GB of fp32 LoRA gradients, reduced once per optimizer step and overlapped with backward, communication should be a small fraction of step time on NVLink-connected GPUs. Scaling efficiency within a node should be high. Measure it rather than assume it.
- Accumulation amortizes it further. With
ACCUM_STEPS=8, seven of every eight backward passes do no communication. - The bottleneck moves to the input pipeline and padding. Once communication is cheap, padding waste (long and short samples in the same batch) and slow tokenization show up. Sorting or bucketing by length, or packing samples, usually buys more than tuning NCCL.
- Across nodes it depends on the network. The same 0.46 GB per GPU per step can dominate on a slow Ethernet link.
To see the overlap directly, capture a few steps with torch.profiler and look at the trace: NCCL all-reduce kernels should run alongside backward compute kernels, not after them. If they don’t overlap, the usual cause is buckets that fill too late, and bucket_cap_mb is the knob to try.
What I verified and what I didn’t
Verified:
- Gemma 4 sizes, parameter counts, layer counts, context lengths, license, release dates, and BF16 inference memory, from the model card, the model overview, the release notes, and the launch post.
- The 12B parameter count, tensor names, global layers without
v_proj, embedding size, and the 65.6M LoRA parameter count, by reading the publishedgoogle/gemma-4-12B-itconfig and safetensors header. - DDP,
DistributedSampler,init_process_group, andtorchrunsignatures and defaults against PyTorch 2.14.0, including the 25 MB default bucket size and theno_sync()warning. - The script itself: I ran a copy on CPU with two
glooprocesses (the only changes were device, backend, and fused AdamW) against a 52M-parameter, randomly initialized model built from the real Gemma 4 12B config and the real tokenizer, with transformers 5.17.0 and PEFT 0.21.0. It loaded, trained, synchronized, logged, checkpointed the adapter, and exited cleanly, and the LoRA parameter checksums matched on both ranks afterward. It also caught the closure-pickling bug described above, and switching the targets toall-linearreproduced the unused-parameter error.
Not verified:
- Any GPU run of Gemma 4 12B, so the memory table is arithmetic and the scaling section is a method, not a result. If you run it,
torch.cuda.max_memory_allocated()and the loggedtok/swill tell you quickly how far off my estimates are. - Whether the larger sizes (26B A4B MoE, 31B) behave the same way under DDP. Their module layouts differ, and the MoE model in particular deserves its own check for unused parameters.
DDP is the least glamorous box in the Megatron diagram, and for fine-tuning it’s usually the right one. Make one replica fit, with LoRA, checkpointing, and sensible sequence lengths. Keep every rank on the same code path. Accumulate with no_sync(). Then measure tokens per second as you add GPUs. The visualizer’s interactive data-parallel tab shows the same loop at toy scale, and if you want to see how this kind of multi-GPU work fits into a delivery platform, the multi-GPU Kubernetes delivery platform case study covers the production side.
Related work
Production ML context
See how this topic connects to production ML systems, infrastructure, and inference.