How to Debug and Diagnose Transformer Training Loss Divergence

Training loss divergence in decoder-only transformers typically stems from misaligned configuration between model context length and training data, missing gradient clipping, or invalid token IDs exceeding vocabulary size, all of which can be systematically detected using the debugging tools in the FareedKhan-dev/train-llm-from-scratch repository.

Training a GPT-style transformer from scratch is notoriously unstable. Loss values may explode to infinity, suddenly spike after thousands of steps, or silently become NaN without warning. This repository implements a minimal transformer architecture, making it an ideal case study for identifying and resolving the root causes of training instability through targeted debugging techniques.

Verify the Loss Computation

The first checkpoint is the loss calculation itself. In src/models/transformer.py, the cross-entropy loss is computed after the final linear projection:

loss = F.cross_entropy(flat_logits, targets)        # src/models/transformer.py L73-L78

Ensure that targets are of type long and match the flattened logits shape. A silent dtype mismatch—such as passing float targets instead of long—will not raise an exception but will produce NaN gradients that propagate through the entire network.

Inspect the Data Pipeline

The data iterator in data_loader/data_loader.py constructs input-target pairs from tokenized HDF5 files:

xb = random_samples[:, :context_length].to(device)   # data_loader.py L55

yb = random_samples[:, 1:context_length+1].to(device) # data_loader.py L56

Three specific configuration mismatches in this pipeline commonly trigger divergence:

  • Context length mismatch: If T_CONTEXT_LENGTH = 16 in your training config but CONTEXT_LENGTH = 512 in the model definition, positional embeddings are only applied to the first 16 tokens. The remaining 496 positions receive no gradient signal, causing the model to output garbage for longer sequences. Align these values in config/config.py before training begins.

  • Out-of-vocabulary tokens: Token IDs in the HDF5 dataset that exceed VOCAB_SIZE cause the embedding lookup to return inf values. Validate that all token IDs are strictly less than VOCAB_SIZE or implement clipping at the data loader level.

  • Batch imbalance: The iterator reshuffles after each epoch, but the counter increment logic may skip the final batch, forcing the model to repeatedly see the same examples. Replace simple increment logic with while counter < n_examples to ensure complete epoch coverage.

Check Optimizer and Learning Rate Configuration

The optimizer definition in scripts/train_transformer.py uses a simple step decay:

optimizer = torch.optim.AdamW(model.parameters(), lr=config['t_lr'])   # train_transformer.py L27-L28

if step == config['t_lr_decay_step']:
    for g in optimizer.param_groups:
        g['lr'] = config['t_lr_decayed']                              # train_transformer.py L16-L20

Learning rate sensitivity is a primary driver of divergence. The default 5e-4 works for small models, but scaling to larger architectures (e.g., 3B parameters) requires rates of 1e-4 or lower.

Lack of warm-up causes large gradient updates in early steps. Implement a linear warm-up for the first 10,000 steps before applying the main cosine schedule to prevent early-step instability.

Monitor and Clip Gradient Magnitudes

The current training loop lacks gradient clipping, allowing exploding gradients to destabilize parameters:

loss.backward()
optimizer.step()

Add gradient norm clipping immediately after backward computation:

loss.backward()
torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=1.0)
optimizer.step()

A max_norm of 1.0 is standard for transformer training. Without this safeguard, gradient norms exceeding 1000.0 will update parameters too aggressively, causing loss spikes.

Enable Automatic Anomaly Detection

When loss becomes NaN, PyTorch’s autograd engine can identify the exact operation responsible. Temporarily wrap the backward pass in anomaly detection mode:

with torch.autograd.detect_anomaly():
    loss.backward()

This will raise an exception at the precise line where invalid values first appear, whether in the attention mechanism, layer normalization, or feed-forward layers.

Validate Causal Masking Implementation

The causal mask is registered as a buffer in src/models/attention.py:

self.register_buffer('tril', torch.tril(torch.ones(context_length, context_length)))  # attention.py L33

If the context_length passed to the Head constructor differs from the actual sequence length fed during training, the mask will be incorrectly sized. This allows the model to attend to future tokens, creating a shortcut that collapses training dynamics. Verify that the same context_length value propagates from Transformer.context_length down to each attention head initialization.

Review Model Initialization Strategy

While the repository uses PyTorch’s default Kaiming initialization, deep stacks with N_BLOCKS = 64 may require more conservative scaling. If divergence persists after checking data and optimization, implement scaled Xavier initialization:

def _reset_parameters(self):
    for p in self.parameters():
        if p.dim() > 1:
            nn.init.xavier_uniform_(p)

Add this method to the Transformer class and invoke it in __init__ to ensure weight magnitudes remain stable through depth.

Implement Real-Time Monitoring

The repository already tracks a running average over AVG_WINDOW = 64 steps in the progress bar. Augment this with explicit gradient norm logging:

grad_norm = torch.norm(
    torch.stack([p.grad.norm() for p in model.parameters() if p.grad is not None])
)
print(f"Grad norm: {grad_norm:.4f}")

For comprehensive visualization, integrate TensorBoard:

from torch.utils.tensorboard import SummaryWriter
writer = SummaryWriter(log_dir="runs/transformer")

# After optimizer.step()

writer.add_scalar("grad_norm", grad_norm, step)
writer.add_scalar("train_loss", loss.item(), step)
writer.add_scalar("lr", optimizer.param_groups[0]["lr"], step)

Practical Debugging Code Examples

Adding Warm-Up and Gradient Clipping

Combine stability techniques in the training loop:

from torch.optim.lr_scheduler import LambdaLR
import math

def warmup_cosine(step):
    warmup_steps = 10_000
    if step < warmup_steps:
        return step / warmup_steps
    return 0.5 * (1 + math.cos(math.pi * (step - warmup_steps) / (config['t_train_steps'] - warmup_steps)))

scheduler = LambdaLR(optimizer, lr_lambda=warmup_cosine)

# Training loop

loss.backward()
torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=1.0)
optimizer.step()
scheduler.step()

Detecting NaN Sources

Pinpoint the exact layer causing divergence:


# Temporary diagnostic code

with torch.autograd.detect_anomaly():
    loss.backward()

Summary

  • Align configuration values: Ensure t_context_length matches context_length in config/config.py to prevent positional embedding mismatch.
  • Clip gradients: Add torch.nn.utils.clip_grad_norm_ with max_norm=1.0 after loss.backward() in scripts/train_transformer.py.
  • Validate token IDs: Confirm all input tokens are less than VOCAB_SIZE to avoid embedding lookup failures.
  • Use warm-up: Implement a linear learning rate warm-up for the first 10,000 steps before cosine decay.
  • Enable anomaly detection: Wrap backward passes with torch.autograd.detect_anomaly() to catch NaN sources immediately.
  • Monitor gradient norms: Log gradient magnitudes each step to detect explosion before loss diverges.

Frequently Asked Questions

Why does transformer loss become NaN immediately after starting training?

Immediate NaN values typically indicate invalid input data or dtype mismatches. Verify that target tensors passed to F.cross_entropy in src/models/transformer.py are long dtype, not float. Additionally, check that token IDs in data_loader/data_loader.py do not exceed VOCAB_SIZE, as out-of-range indices return inf from the embedding layer.

How do I fix exploding gradients in a from-scratch GPT implementation?

Apply gradient clipping using torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=1.0) immediately after loss.backward() and before optimizer.step(). If gradients still explode, reduce the learning rate from 5e-4 to 1e-4 and implement a 10,000-step linear warm-up scheduler to prevent large early updates.

What causes sudden loss spikes after thousands of training steps?

Sudden divergence often occurs when the learning rate decays too abruptly or when the causal mask configuration drifts. In scripts/train_transformer.py, verify that t_lr_decay_step occurs after warm-up completion. Check that context_length in src/models/attention.py matches the sequence length produced by data_loader/data_loader.py to prevent attention to future tokens.

How can I identify which specific layer is causing training instability?

Enable PyTorch’s anomaly detection mode by wrapping the backward pass in with torch.autograd.detect_anomaly():. This will raise an exception at the exact tensor operation producing NaN or infinite values. Additionally, log gradient norms per layer using p.grad.norm() for each parameter to isolate whether the attention heads, feed-forward layers, or embeddings are generating unstable gradients.

Have a question about this repo?

These articles cover the highlights, but your codebase questions are specific. Give your agent direct access to the source. Share this with your agent to get started:

Share the following with your agent to get started:
curl -s "https://instagit.com/install.md"

Works with
Claude Codex Cursor VS Code OpenClaw Any MCP Client

Maintain an open-source project? Get it listed too →