# How to Debug and Diagnose Transformer Training Loss Divergence

> Debug and diagnose transformer training loss divergence. Discover common causes like context length misalignment, missing gradient clipping, or invalid tokens. Utilize debugging tools from train-llm-from-scratch to fix issues.

- Repository: [Fareed Khan/train-llm-from-scratch](https://github.com/FareedKhan-dev/train-llm-from-scratch)
- Tags: how-to-guide
- Published: 2026-05-31

---

**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`](https://github.com/FareedKhan-dev/train-llm-from-scratch/blob/main/src/models/transformer.py), the cross-entropy loss is computed after the final linear projection:

```python
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`](https://github.com/FareedKhan-dev/train-llm-from-scratch/blob/main/data_loader/data_loader.py) constructs input-target pairs from tokenized HDF5 files:

```python
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`](https://github.com/FareedKhan-dev/train-llm-from-scratch/blob/main/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`](https://github.com/FareedKhan-dev/train-llm-from-scratch/blob/main/scripts/train_transformer.py) uses a simple step decay:

```python
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:

```python
loss.backward()
optimizer.step()

```

Add gradient norm clipping immediately after backward computation:

```python
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:

```python
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`](https://github.com/FareedKhan-dev/train-llm-from-scratch/blob/main/src/models/attention.py):

```python
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:

```python
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:

```python
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:

```python
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:

```python
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:

```python

# 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`](https://github.com/FareedKhan-dev/train-llm-from-scratch/blob/main/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`](https://github.com/FareedKhan-dev/train-llm-from-scratch/blob/main/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`](https://github.com/FareedKhan-dev/train-llm-from-scratch/blob/main/src/models/transformer.py) are `long` dtype, not `float`. Additionally, check that token IDs in [`data_loader/data_loader.py`](https://github.com/FareedKhan-dev/train-llm-from-scratch/blob/main/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`](https://github.com/FareedKhan-dev/train-llm-from-scratch/blob/main/scripts/train_transformer.py), verify that `t_lr_decay_step` occurs after warm-up completion. Check that `context_length` in [`src/models/attention.py`](https://github.com/FareedKhan-dev/train-llm-from-scratch/blob/main/src/models/attention.py) matches the sequence length produced by [`data_loader/data_loader.py`](https://github.com/FareedKhan-dev/train-llm-from-scratch/blob/main/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.