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:
config/config.py– Add the warmup duration parameterscripts/train_transformer.py– Wrap the optimizer with the scheduler- Training loop – Call
scheduler.step()after eachoptimizer.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_stepsparameter toconfig/config.pyto control warmup duration - Creating a
LambdaLRscheduler inscripts/train_transformer.pythat scales the learning rate linearly from0to1over the warmup period - Calling
scheduler.step()immediately afteroptimizer.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:
curl -s "https://instagit.com/install.md" Maintain an open-source project? Get it listed too →