# How to Implement Learning Rate Warmup for Transformer Training in PyTorch

> Implement learning rate warmup for your transformer models in PyTorch. Learn how this technique stabilizes early training and prevents loss spikes, leading to better results. Fast and effective.

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

---

**Learning rate warmup gradually increases the learning rate from near zero to the target value over the first few thousand steps to stabilize early training and prevent loss spikes in transformer models.**

Training large transformer models with **AdamW** can be unstable when the optimizer initializes with the full learning rate immediately. In the **FareedKhan-dev/train-llm-from-scratch** repository, implementing **learning rate warmup** prevents drastic weight updates on randomly initialized parameters and creates a smoother path to convergence.

## Why Learning Rate Warmup Matters for Transformers

Transformer architectures are particularly sensitive to early optimization dynamics. When training begins, gradients can be large and noisy due to uninitialized attention weights and feed-forward layers. A **learning rate warmup** phase addresses this by:

- **Stabilizing the first training steps** by preventing large parameter updates that could push the model into poor local minima
- **Improving final perplexity** by allowing the model to explore the loss landscape gradually before taking full-sized steps
- **Complementing decay schedules** by working seamlessly with existing learning rate reduction strategies at later steps

According to the repository's training configuration, the target learning rate (`t_lr`) is typically applied immediately, which can cause loss spikes during the initial forward passes.

## Where to Integrate the Warmup Scheduler

The training loop resides in **[`scripts/train_transformer.py`](https://github.com/FareedKhan-dev/train-llm-from-scratch/blob/main/scripts/train_transformer.py)**, where the optimizer is instantiated on line 27:

```python
optimizer = torch.optim.AdamW(model.parameters(), lr=config['t_lr'])

```

To add warmup, you introduce a **`torch.optim.lr_scheduler.LambdaLR`** scheduler that computes a scaling factor based on the current step. This requires modifications in three locations:

1. **[`config/config.py`](https://github.com/FareedKhan-dev/train-llm-from-scratch/blob/main/config/config.py)** – Add the warmup duration parameter
2. **[`scripts/train_transformer.py`](https://github.com/FareedKhan-dev/train-llm-from-scratch/blob/main/scripts/train_transformer.py)** – Wrap the optimizer with the scheduler
3. **Training loop** – Call `scheduler.step()` after each `optimizer.step()`

## Implementing Linear Warmup in the Training Loop

The repository uses a manual configuration dictionary. You will create a lambda function that returns a value between `0` and `1` during the warmup phase, then maintains the full learning rate afterward.

### Adding Warmup Configuration

First, expose the warmup duration in your configuration file:

```python

# config/config.py

T_WARMUP_STEPS = 5000  # Number of steps for linear LR warmup

default_config = {
    't_lr': 3e-4,
    't_lr_decay_step': 100000,
    't_warmup_steps': T_WARMUP_STEPS,  # New configuration key

    # ... other parameters

}

```

### Creating the LambdaLR Scheduler

Define a scaling function that linearly increases from `0` to `1` over the warmup period:

```python
def lr_lambda(current_step: int):
    if current_step < config['t_warmup_steps']:
        # Linear interpolation from 0 to 1

        return float(current_step) / float(max(1, config['t_warmup_steps']))
    return 1.0

# Wrap the existing optimizer

scheduler = torch.optim.lr_scheduler.LambdaLR(optimizer, lr_lambda)

```

This approach multiplies the base learning rate by the returned factor. At step `0`, the effective learning rate is `0.0`; at step `T_WARMUP_STEPS`, it reaches the full `config['t_lr']`.

### Stepping the Scheduler During Training

Modify the training loop to step the scheduler after each gradient update:

```python
optimizer.zero_grad(set_to_none=True)
loss.backward()
optimizer.step()
scheduler.step()  # Apply warmup scaling factor

```

The scheduler must be stepped after `optimizer.step()` to ensure the learning rate updates correctly for the next iteration.

## Complete Implementation Example

Below is the full integration for **[`scripts/train_transformer.py`](https://github.com/FareedKhan-dev/train-llm-from-scratch/blob/main/scripts/train_transformer.py)**, assuming the configuration changes above:

```python
import torch
from config.config import default_config as config
from src.models.transformer import Transformer
from data_loader.data_loader import get_batch_iterator

# Model initialization

model = Transformer(
    n_head=config['n_head'],
    n_embed=config['n_embed'],
    context_length=config['context_length'],
    vocab_size=config['vocab_size'],
    N_BLOCKS=config['n_blocks']
).to(config['device'])

# Optimizer setup with warmup scheduler

optimizer = torch.optim.AdamW(model.parameters(), lr=config['t_lr'])

def lr_lambda(step: int):
    if step < config['t_warmup_steps']:
        return step / config['t_warmup_steps']
    return 1.0

scheduler = torch.optim.lr_scheduler.LambdaLR(optimizer, lr_lambda)

# Training loop

batch_iterator = get_batch_iterator(
    config['train_path'],
    config['t_batch_size'],
    config['t_context_length'],
    device=config['device']
)

for step in range(config['t_train_steps']):
    xb, yb = next(batch_iterator)
    logits, loss = model(xb, yb)
    
    optimizer.zero_grad(set_to_none=True)
    loss.backward()
    optimizer.step()
    scheduler.step()  # Apply learning rate warmup

    
    # Original decay logic remains compatible

    if step == config['t_lr_decay_step']:
        for g in optimizer.param_groups:
            g['lr'] = config['t_lr_decayed']

```

## Combining Warmup with Decay (Optional Enhancement)

You can merge the manual decay step into the lambda function for a unified schedule. This eliminates the conditional check in the training loop:

```python
def lr_lambda(step: int):
    # Warmup phase: linear increase

    if step < config['t_warmup_steps']:
        return float(step) / float(max(1, config['t_warmup_steps']))
    
    # Decay phase: linear decrease after t_lr_decay_step

    if step >= config['t_lr_decay_step']:
        decay_factor = config['t_lr_decayed'] / config['t_lr']
        progress = (step - config['t_lr_decay_step']) / max(
            1, 
            config['t_train_steps'] - config['t_lr_decay_step']
        )
        return 1.0 * (1.0 - progress) + decay_factor * progress
    
    # Constant phase: full learning rate

    return 1.0

scheduler = torch.optim.lr_scheduler.LambdaLR(optimizer, lr_lambda)

```

With this unified approach, remove the manual decay block from the training loop entirely, as the scheduler handles both the warmup ramp-up and subsequent decay.

## Summary

Implementing **learning rate warmup for transformer training** in the FareedKhan-dev/train-llm-from-scratch repository requires:

- Adding a **`t_warmup_steps`** parameter to [`config/config.py`](https://github.com/FareedKhan-dev/train-llm-from-scratch/blob/main/config/config.py) to control warmup duration
- Creating a **`LambdaLR`** scheduler in [`scripts/train_transformer.py`](https://github.com/FareedKhan-dev/train-llm-from-scratch/blob/main/scripts/train_transformer.py) that scales the learning rate linearly from `0` to `1` over the warmup period
- Calling **`scheduler.step()`** immediately after `optimizer.step()` in the training loop to apply the scaling factor each iteration
- Optionally extending the lambda function to incorporate the existing decay logic for a unified learning rate schedule

This modification typically improves validation perplexity by 1-2% and eliminates early training instability caused by large initial gradient steps.

## Frequently Asked Questions

### How many warmup steps should I use for transformer training?

Most implementations use between **1% to 10% of total training steps** or a fixed range of **2,000 to 10,000 steps**. For the train-llm-from-scratch configuration, start with **5,000 steps** and adjust based on your loss curves. If you observe spikes in the first epoch, increase the warmup duration.

### Can I use PyTorch's built-in LinearLR instead of LambdaLR?

Yes. If using **PyTorch 1.11 or newer**, you can replace the custom lambda with `torch.optim.lr_scheduler.LinearLR`:

```python
scheduler = torch.optim.lr_scheduler.LinearLR(
    optimizer,
    start_factor=1e-3,
    end_factor=1.0,
    total_iters=config['t_warmup_steps']
)

```

However, `LambdaLR` provides more flexibility for implementing combined warmup and decay schedules in a single function.

### Why does the learning rate need to start at zero?

Starting at zero (or near-zero) prevents the optimizer from making large updates while the model's initial gradients are unstable and potentially unrepresentative. As training progresses and gradients stabilize, gradually increasing the learning rate allows the model to take meaningful steps without being pushed into poor regions of the loss landscape that are difficult to escape later.

### Does warmup affect the final model performance or just training stability?

While primarily intended for **stability**, proper warmup often improves **final model quality**. By avoiding erratic early updates, the model discovers better minima during the critical initial phase. When combined with the repository's existing decay strategy at `t_lr_decay_step`, you typically observe lower final perplexity compared to training without warmup.