Skip to content
InteractiveLive

Training a transformer on eight GPUs, then serving it

One concrete configuration, followed from start to finish: a 1.27B-parameter model with Llama 2 7B's layer shape, trained on a single 8-GPU node split 2-way tensor × 2-way pipeline × 2-way data parallel, then post-trained and deployed.

Published
Topics
  • Distributed training
  • Parallelism
  • LLM training
  • Inference
  • Interactive

One model, one machine, every number

Explanations of distributed training usually stay abstract. This one fixes a single setup and computes everything for it: 5 layers with Llama 2 7B's exact layer shape (hidden size 4,096, 32 heads, a SwiGLU MLP 11,008 wide, a 32,000-token vocabulary), 1.27 billion parameters, 2,048-token sequences and a global batch of 32,768 tokens per step.

The model is shallow but every layer is full size, so the per-layer shapes, bytes, FLOPs and communication are what you would see in a real 7B-class model. One colour scheme marks each kind of parallelism throughout.

What it covers

  • The lifecycle, from data through pretraining, post-training and evaluation to deployment
  • One training step on one GPU: forward, loss, backward and the AdamW update, with mixed precision at 16 bytes per parameter
  • Follow the matrices: a tiny model with RMSNorm, RoPE, causal attention and SwiGLU, run and trained live in your browser
  • Why one GPU runs out: the memory wall and the compute wall
  • Data parallelism, and ZeRO stages 1 to 3 and FSDP
  • Tensor and sequence parallelism, Megatron-style
  • Pipeline parallelism with micro-batches and the 1F1B schedule
  • Context parallelism for long sequences, and expert parallelism for mixture-of-experts
  • All together: an animated full step on a 2 × 2 × 2 device mesh
  • Keeping a long run alive: checkpoints, failures and loss spikes
  • Post-training: supervised fine-tuning, RLHF, DPO and GRPO
  • Deployment: merging shards, quantization, the KV cache, batching and speculative decoding

Who it is for

Anyone who understands a single transformer and wants to know what actually happens when it is trained across many GPUs: which tensors are split, who talks to whom and how often, and where the memory goes. Numbers are rounded estimates for building intuition, and the reading list names the papers behind each technique.

To see the single block that all of this splits up, start with the neuron-by-neuron explainer.

Questions

What is the difference between data, tensor and pipeline parallelism?
Data parallelism gives every GPU a full copy of the model and a different slice of the batch. Tensor parallelism splits each weight matrix across GPUs, so they cooperate on every layer. Pipeline parallelism puts different groups of layers on different GPUs and streams micro-batches through them. Large runs combine all three.
How much GPU memory does training a model take?
With bf16 matmuls and AdamW, about 16 bytes per parameter for weights, gradients and optimizer state, before any activations: roughly 20 GB for a 1.27B model and over 1 TB for a 70B model. That is why model state has to be sharded.
What is the difference between ZeRO and FSDP?
Both shard model state across data-parallel GPUs. ZeRO stage 1 shards optimizer state, stage 2 also shards gradients, and stage 3 also shards the weights themselves. FSDP is PyTorch's implementation of the stage 3 approach.
How long does a training step take on 8 GPUs?
For the example (1.27B parameters, 32,768 tokens per step), about 240 TFLOP per step. Eight H100s at around 40% utilization finish it in roughly 0.08 seconds, about 430,000 tokens per second.

Credits

Training a transformer on eight GPUs stands on open-source work and published research. Thank you to everyone behind it.

More from Sarvabhaum

Web appLive

attnlab

A set of labs, walked in order, that run real TransformerLens models behind a web page. Type a prompt, pick a model, and see how it is tokenized, where every head attends, and when the model knows its answer.

  • Mechanistic interpretability
  • TransformerLens
  • Attention
InteractiveLive

One transformer block, neuron by neuron

A tiny real transformer (d = 8, 2 heads, an 8 → 32 → 8 MLP, 4 blocks) runs on the sentence you type. Every circle holds the number it actually computed, and hovering one shows the weighted lines that feed it.

  • Transformers
  • Neural networks
  • Attention
InteractiveLive

A transformer block, one matrix at a time

Follow four tokens through a single transformer block from the residual stream's point of view. Every matrix is shown with real numbers, computed live, so you can trace one token's row from embedding to next-token probabilities.

  • Transformers
  • Residual stream
  • Attention