Implementing Learning Rate Schedulers for LLM Pretraining: Warm-Up and Cosine Decay in PyTorch

The rasbt/LLMs-from-scratch repository implements a two-phase learning rate scheduler that combines linear warm-up with cosine annealing directly inside the training loop, eliminating external dependencies while ensuring stable convergence for transformer models.

Implementing learning rate schedulers for LLM pretraining is essential for preventing early-training instabilities and maximizing final model performance. Rather than relying on external scheduling libraries, the rasbt/LLMs-from-scratch repository embeds a lightweight warm-up and cosine decay mechanism directly into the training orchestration code, providing full transparency and control over the optimization trajectory.

How the Learning Rate Scheduler Works

The scheduler operates in two distinct phases that mirror the optimization strategies used in modern large language models like GPT-3. Implemented in pkg/llms_from_scratch/appendix_d.py, the logic automatically adjusts the optimizer's learning rate at every global step based on the current training progress.

Linear Warm-Up Phase

During the initial warmup_steps, the learning rate increases linearly from an initial_lr to the peak learning rate defined in the optimizer configuration. According to lines 36-38 in appendix_d.py, the increment per step is calculated as:

lr_increment = (peak_lr - initial_lr) / warmup_steps
current_lr = initial_lr + global_step * lr_increment

This gradual ramp prevents large gradient spikes during the early stages of training when the model weights are far from optimal. The peak learning rate is extracted dynamically from optimizer.param_groups[0]["lr"], ensuring compatibility with any PyTorch optimizer that respects the standard parameter group interface.

Cosine Annealing Phase

Once the warm-up phase completes (global_step >= warmup_steps), the scheduler transitions to a cosine decay curve that smoothly reduces the learning rate from the peak down to a specified min_lr. As implemented in lines 50-54 of appendix_d.py, the calculation follows:

progress = (global_step - warmup_steps) / (total_steps - warmup_steps)
current_lr = min_lr + (peak_lr - min_lr) * 0.5 * (1 + cos(pi * progress))

This cosine annealing approach often yields better final perplexity scores compared to step-wise or exponential decay schedules, as it allows for fine-grained weight updates during the later stages of training when the model is approaching convergence.

Integration with the Training Loop

The train_model function in pkg/llms_from_scratch/appendix_d.py serves as the central orchestrator, handling not only the forward and backward passes but also the learning rate updates and gradient management. The scheduler writes the computed learning rate back to every parameter group via param_group["lr"] = current_lr, making it transparent to the underlying optimizer.

Additionally, the implementation includes gradient clipping (lines 64-70) to prevent exploding gradients. Note that the corrected version uses >= rather than > for the clipping threshold comparison, ensuring proper handling of boundary cases during the warm-up phase.

Implementation Examples

Basic Usage with train_model

To utilize the scheduler in a standard training script such as ch05/01_main-chapter-code/gpt_train.py, pass the warm-up and learning rate bounds directly to the training function:

import torch
from pkg.llms_from_scratch.appendix_d import train_model

# Initialize model and optimizer

model = MyTransformer(vocab_size=50257, d_model=768, n_heads=12, ...)
optimizer = torch.optim.AdamW(model.parameters(), lr=5e-4)  # This becomes peak_lr

# Configure scheduler parameters

warmup_steps = 500
initial_lr = 3e-5
min_lr = 1e-6

# Execute training with automatic LR scheduling

train_losses, val_losses, tokens_seen, track_lrs = train_model(
    model=model,
    train_loader=train_loader,
    val_loader=val_loader,
    optimizer=optimizer,
    device="cuda",
    n_epochs=10,
    eval_freq=1000,
    warmup_steps=warmup_steps,
    initial_lr=initial_lr,
    min_lr=min_lr,
    orig_book_version=False,  # Use corrected gradient clipping

)

The function returns track_lrs, a list containing the learning rate value at every optimization step, which enables precise monitoring and debugging of the schedule behavior.

Hyperparameter Search Over Warm-Up Steps

For systematic exploration of optimal warm-up durations, the repository provides ch05/05_bonus_hparam_tuning/hparam_search.py. This utility accepts a grid configuration to test different warmup_iters values:

from ch05_05_bonus_hparam_tuning.hparam_search import run_hparam_search

hparam_grid = {
    "warmup_iters": [100, 250, 500, 1000],
    "learning_rate": [1e-4, 5e-4, 1e-3],
}

run_hparam_search(
    train_loader=train_loader,
    val_loader=val_loader,
    model_factory=lambda: MyTransformerModel(),
    optimizer_factory=lambda model: torch.optim.AdamW(model.parameters(), lr=5e-4),
    hparam_config=hparam_grid,
    max_epochs=3,
)

Each trial automatically instantiates the scheduler with the specified warmup_iters, allowing you to identify the configuration that minimizes validation loss for your specific dataset and model architecture.

Monitoring the Learning Rate Schedule

Visualizing the learning rate progression helps verify that the warm-up and decay phases are operating as expected:

import matplotlib.pyplot as plt

# track_lrs is returned by train_model

plt.plot(track_lrs)
plt.title("Learning Rate Schedule: Warm-Up + Cosine Decay")
plt.xlabel("Training Step")
plt.ylabel("Learning Rate")
plt.axvline(x=warmup_steps, color='r', linestyle='--', label='End of Warm-Up')
plt.legend()
plt.show()

This inspection is particularly valuable when scaling to distributed training scenarios, such as those implemented in ch05/10_llm-training-speed/02_opt_multi_gpu_ddp.py, where ensuring consistent learning rate updates across processes is critical.

Key Files and Source Locations

Summary

  • The rasbt/LLMs-from-scratch repository implements a linear warm-up with cosine annealing schedule directly in pkg/llms_from_scratch/appendix_d.py, requiring no external scheduler classes.
  • The scheduler automatically extracts the peak learning rate from the optimizer's initial configuration and interpolates between initial_lr, peak_lr, and min_lr based on global_step.
  • Gradient clipping is integrated into the same training loop (lines 64-70), with the corrected version using proper boundary checking to ensure numerical stability.
  • The train_model function returns a track_lrs list that enables precise monitoring and visualization of the learning rate trajectory.
  • Hyperparameter search capabilities in hparam_search.py allow systematic exploration of warm-up durations to optimize training dynamics for specific hardware and datasets.

Frequently Asked Questions

Why combine warm-up with cosine annealing for LLM pretraining?

Warm-up stabilizes early training by preventing large gradient steps when the model is initialized with random weights, while cosine annealing provides a smooth decay that allows the optimizer to settle into flat minima rather than bouncing between sharp ones. This combination is empirically shown to produce lower final perplexity than constant or step-decay schedules in transformer architectures.

How do I adjust the warm-up duration for different model sizes?

Modify the warmup_steps parameter passed to train_model or the warmup_iters key in the hyperparameter search grid. As a rule of thumb documented in the repository's training scripts, larger models typically benefit from longer warm-up periods (500-2000 steps) to accommodate more complex loss landscapes, while smaller models may converge with as few as 100-250 warm-up steps.

Can I use this scheduler with optimizers other than AdamW?

Yes. The scheduler is optimizer-agnostic because it writes the computed learning rate directly to param_group["lr"] for every parameter group in the optimizer. This approach works with any PyTorch optimizer that respects the learning rate field in its parameter groups, including SGD, Adam, and AdamW variants.

How does gradient clipping interact with the learning rate schedule?

Gradient clipping (applied at lines 64-70 in appendix_d.py) operates after the learning rate update, ensuring that the scaled gradients remain within a norm threshold regardless of the current learning rate magnitude. This sequencing prevents the large parameter updates that can occur during the transition from warm-up to cosine phases, maintaining training stability throughout the optimization process.

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 →