How to Implement Learning Rate Warmup for Transformer Training in PyTorch

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, where the optimizer is instantiated on line 27:

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 – Add the warmup duration parameter
  2. 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:


# 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:

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:

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, assuming the configuration changes above:

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:

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 to control warmup duration
  • Creating a LambdaLR scheduler in 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:

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.

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 →