Skip to content

Latest commit

 

History

5 Commits

Folders and files

NameName
Last commit message
Last commit date
 
 
 
 
 
 
 
 
 
 
 
 
 
 

Repository files navigation

Distributed Training of a Llama-style Transformer for Text Generation

Multi-GPU distributed training of a 128M-parameter decoder-only transformer (Llama architecture) on the BookCorpus dataset using PyTorch DDP on RunPod.

Architecture

Component Implementation
Position Encoding Rotary Position Embeddings (RoPE)
Normalization RMSNorm (pre-norm)
Attention Grouped Query Attention (GQA) — 12 query heads, 4 KV heads
Feed-Forward SwiGLU — 3 linear layers with SiLU gating
Bias None
Weight Tying Embedding ↔ lm_head
Attention Kernel SDPA (FlashAttention2)

Config: d_model=768, n_layers=12, n_heads=12, n_kv_heads=4, d_ff=2560, max_seq_len=512 — 128.3M parameters


Quick Start (RunPod — 4x RTX 5090 Spot)

1. Create Pod

  • runpod.io → GPU Pod → 4x RTX 5090 (spot, $0.64/GPU/hr)
  • Template: RunPod PyTorch
  • Volume: attach at /workspace

2. Run Everything

cd /workspace
git clone <repo-url> horovod-transformer && cd horovod-transformer
bash scripts/run_all.sh

That's it. run_all.sh handles setup (~2 min), preprocessing (~30 min), and training (~3.5 hrs).

Total cost: ~$10-11 out of $12 budget.

3. If Spot Gets Interrupted

cd /workspace/horovod-transformer
bash scripts/train.sh configs/config.yaml 4 --resume checkpoints/emergency_checkpoint.pt

SIGTERM handler auto-saves a checkpoint when the spot instance is about to die.

4. Generate Text

python -m src.generate \
  --checkpoint checkpoints/best_model.pt \
  --prompt "Once upon a time" \
  --max_new_tokens 200

# Interactive mode
python -m src.generate --checkpoint checkpoints/best_model.pt --interactive

Project Structure

src/
├── models/transformer_lm.py      # Llama-style model (RoPE, RMSNorm, SwiGLU, GQA)
├── train.py                       # DDP training + SIGTERM handler for spot
├── data.py                        # DistributedSampler data loading
├── generate.py                    # Text generation (top-k, top-p, repetition penalty)
├── metrics.py                     # Perplexity, accuracy
└── utils/
    ├── distributed.py             # PyTorch DDP utilities (torchrun)
    ├── config.py                  # YAML config
    ├── logging.py                 # Rank-aware logging + TensorBoard
    └── recorder.py                # Structured CSV metrics

configs/
├── config.yaml                    # 4x RTX 5090 training config
└── config_small.yaml              # Debug config (CPU-friendly)

scripts/
├── run_all.sh                     # One-shot: setup + preprocess + train
├── setup_droplet.sh               # Pod environment setup
├── preprocess.sh                  # Data preprocessing
└── train.sh                       # torchrun launcher

Training Details

Parameter Value
Framework PyTorch DDP via torchrun
Optimizer AdamW (fused, betas=0.9/0.95)
LR Schedule Cosine decay 3e-4 → 3e-5, 2000 warmup
Batch size 32/GPU × 3 accum × 4 GPUs = 384
Precision bf16
Gradient sync no_sync() on micro-steps
Spot safety SIGTERM → emergency checkpoint

Monitoring

# TensorBoard (SSH tunnel from local machine)
ssh -L 6006:localhost:6006 <pod>
tensorboard --logdir runs/

WandB logs automatically if WANDB_API_KEY is set.


Cost

Phase Time Cost
Setup ~2 min $0.09
Preprocessing ~30 min $1.28
Training (3 epochs) ~3.5 hrs $8.96
Total ~4 hrs ~$10.33

About

Distributed training of a Llama-style 128M-param transformer on BookCorpus using PyTorch DDP on RunPod (4x RTX 5090)

Resources

Stars

0 stars

Watchers

0 watching

Forks

Releases

Packages

Contributors

Languages