Most models are not trained once. New documents arrive, a product adds a language, a domain team collects a fresh corpus, and the question becomes whether to retrain from scratch on everything or continue from the model you already have. Incremental training means continuing: starting from an existing checkpoint and training further on new data, so the compute you already spent is reused. Done well it costs a small fraction of a full run. Done carelessly it causes catastrophic forgetting, loss spikes, or a run that silently restarts its data from the beginning.

This article treats incremental training as a GPU engineering problem as much as a modelling one. It covers what a resumable checkpoint must contain and how big it is, and how to re-shape the learning-rate schedule. It then deals with mixing in old data to prevent forgetting and restoring sharded state on a different number of GPUs, with code using PyTorch distributed checkpointing. A worked example compares the compute of continuing against retraining, and the article ends with adapters as the lightweight alternative, failure modes and a checklist. Memory budgets for training itself are covered in full fine-tuning on GPUs.

Three kinds of incremental training

RegimeWhat changesTypical tool
ResumeNothing; the same run continues after a crash or pre-emptionExact checkpoint restore
Continual pre-trainingNew unlabelled data, often a shifted distributionFull-parameter training, new LR schedule, replay
Incremental fine-tuningNew task or instruction data on a tuned modelFull fine-tune or a new adapter per increment

The regimes share machinery but differ in risk. Resuming must be bit-for-bit faithful: the goal is that the run behaves as if it had never stopped. Continual pre-training deliberately changes the data, so the danger is forgetting what the old data taught. Incremental fine-tuning is smaller and is often better served by adapters, which leave the base weights untouched.

What a resumable checkpoint holds

A checkpoint for incremental training is not the same as a model export. An export holds the weights in the precision you serve, usually BF16 at two bytes per parameter. A resumable checkpoint holds everything the optimiser and data pipeline need to carry on:

StateSize for a 7B modelWhy it matters
FP32 master weights28 GB (4 bytes per param)Mixed-precision training updates these, not the BF16 copy
Adam first moment m28 GBMomentum; losing it makes the first steps behave like a cold start
Adam second moment v28 GBPer-parameter step sizes; resetting it causes loss spikes
Scheduler and step countbytesWhere the learning rate is and what it does next
Loss scaler (FP16 only)bytesCurrent scale factor; BF16 runs usually do not need one
Data loader positionkilobytesWhich samples were seen; otherwise data repeats or is skipped
RNG states, per rankkilobytesDropout masks and shuffles stay reproducible

Gradients are not saved: they are recomputed on the next step. So a resumable 7B checkpoint is about 84 GB against a 14 GB export, six times larger. Across 64 GPUs with sharded optimiser state, each rank writes about 1.3 GB, and at a sustained aggregate write rate of a few gigabytes per second to shared storage, a save takes tens of seconds. That sets how often you can afford to checkpoint. The mixed-precision guide explains why master weights stay in FP32.

The learning-rate schedule

The most common mistake in continual pre-training is the learning rate. A run trained with cosine decay ends with a learning rate near zero. Continue at that rate and the model barely learns the new data. Jump straight back to the peak and the first few hundred steps produce a loss spike that damages what the model already knew. Ibrahim and colleagues studied this in 2024 at 405 million and 10 billion parameters. They found that re-warming the learning rate, re-decaying it, and replaying a portion of the previous data together matched retraining from scratch on all the data, measured by final loss and average benchmark score, for a fraction of the compute.

import math

def continual_lr(step, total, peak, warmup, floor_ratio=0.1):
    """Fresh linear warmup to `peak`, then cosine decay to floor_ratio * peak."""
    if step < warmup:
        return peak * (step + 1) / warmup
    progress = (step - warmup) / max(1, total - warmup)
    floor = peak * floor_ratio
    return floor + 0.5 * (peak - floor) * (1 + math.cos(math.pi * progress))

Choose the re-warm peak below the original peak when the new data is close to the old, and nearer it when the shift is large and you need the model to move further. Warmup length matters more than you would expect: a longer re-warm gives the restored Adam statistics time to adapt to the new gradient distribution. The same paper proposes schedules that are not tied to a fixed token budget. A practical form is warmup, a long constant phase, then a short decay: keep checkpoints from the constant phase and branch each new increment from there rather than from the fully decayed endpoint.

Continuing a finished run on new data: learning rate and data mixLRtokensoriginal run: warmup, cosine decaycheckpointre-warm, then re-decayoriginal corpusnew datareplayRestore exact stateweights, Adam m and v, RNGNew schedulefresh warmup and decayMixed loadernew data plus old samples
The original run decays to a low learning rate. Continuing restores full optimiser state, starts a fresh warmup and decay, and feeds a mix of new data with a slice of replayed old data.

Data: replay, tokenizer and loader state

Replay means mixing samples from the original corpus into the new data stream. Its purpose is to keep the gradient signal of the old distribution present so the model does not drift away from it. How much you need depends on how far the new data shifts; the Ibrahim study used larger fractions for its stronger English-to-German shift than for its English-to-English one. Start with a small fraction for a mild shift, measure loss on a held-out slice of the old data, and increase replay if it rises.

Four data rules are easy to break. Keep the tokenizer identical, because new tokens require resizing the embedding and output layers and initialising the new rows sensibly, which is a separate project. Deduplicate the new data against the old, or you will over-train on repeats. Keep evaluation sets out of both. And make the data loader stateful, so a resume within an increment continues where it stopped instead of starting the epoch again.

Saving and restoring sharded state

PyTorch's distributed checkpointing, torch.distributed.checkpoint or DCP, saves each rank's shard in parallel and can load into a different sharding layout, which matters because incremental runs often get a different GPU allocation than the original. The state-dict helpers produce model and optimiser state in a form DCP understands for plain, DDP and FSDP models:

import torch
import torch.distributed.checkpoint as dcp
from torch.distributed.checkpoint.state_dict import get_state_dict, set_state_dict
from torchdata.stateful_dataloader import StatefulDataLoader

def save(path, model, optim, sched, loader, step):
    model_sd, optim_sd = get_state_dict(model, optim)
    state = {
        "model": model_sd, "optim": optim_sd,
        "sched": sched.state_dict(), "loader": loader.state_dict(),
        "step": step, "rng": torch.cuda.get_rng_state(),   # per-rank: see note below
    }
    dcp.save(state, checkpoint_id=path)

def load_for_continual(path, model, optim):
    """Restore weights and Adam moments only; schedule and data start fresh."""
    model_sd, optim_sd = get_state_dict(model, optim)       # templates in the current layout
    state = {"model": model_sd, "optim": optim_sd}
    dcp.load(state, checkpoint_id=path)                     # reshards if the world size changed
    set_state_dict(model, optim, model_state_dict=state["model"],
                   optim_state_dict=state["optim"])

loader = StatefulDataLoader(mixed_dataset, batch_size=8, num_workers=4)

Two functions, deliberately. A crash resume restores everything, including the scheduler, loader position and RNG. A continual increment restores only weights and optimiser moments, builds a new scheduler with the re-warm schedule, and builds a new loader over the mixed dataset. RNG state differs per rank, so store it per rank rather than in a single shared key; the code shows where it goes, not a complete per-rank solution. The optimiser must have the same parameter groups as when it was saved; changing which parameters are trained, or freezing layers, breaks the mapping. See FSDP for how sharded state is laid out.

Worked example: continue or retrain?

A team has a 7B model pre-trained on 1 trillion tokens and receives 50 billion tokens of new domain text. Training compute is about 6 times parameters times tokens. Retraining on everything costs 6 x 7e9 x 1.05e12, about 4.4e22 FLOPs. Continuing on the new data plus 10 billion replayed tokens costs 6 x 7e9 x 6e10, about 2.5e21 FLOPs, roughly 18 times less.

On H100 GPUs sustaining around 400 TFLOPS of useful BF16 work each, about 40 percent of the dense peak, retraining needs roughly 30,000 GPU-hours. Continuing needs roughly 1,750: two days on 32 GPUs rather than about 40 days. Checkpoint cost changes the plan only slightly. Saving the 84 GB state every 30 minutes across 32 GPUs adds about one percent of overhead if each save takes 20 seconds of blocking time; asynchronous saving with dcp.async_save hides most of it. The team evaluates before and after on three suites: held-out new-domain text, held-out original-distribution text, and its standard capability benchmarks. It accepts the increment only if new-domain loss falls and the other two stay within noise. For deeper per-step cost arithmetic, see fine-tuning maths.

Adapters as increments

When the increment is a task, not a corpus, training a new low-rank adapter per increment is often better. The base weights never change, so nothing is forgotten by construction. Each increment is a few hundred megabytes, rollback means unloading an adapter, and many increments can be served from one base with multi-LoRA serving. The trade-off is capacity: adapters learn new behaviour and formats well but absorb large amounts of new knowledge less well than full-parameter training. When adapters pile up, periodically merge the stable ones into a new base with a proper continual-training run.

Failure modes

  • Optimiser state dropped. Loading only model weights silently cold-starts Adam and produces a spike. Check that moments were restored by logging the norm of the first step's update.
  • Loader restarts at zero. A non-stateful loader repeats early data after every resume. Save and restore its state, and log the first sample IDs after resume.
  • Forgetting discovered late. Only evaluating the new domain hides regressions. Always track held-out loss on the old distribution.
  • Re-warm too hot. Peak too high or warmup too short causes a lasting loss jump. Lower the peak or lengthen warmup.
  • Parameter group mismatch. Adding, removing or freezing parameters makes the saved optimiser state fail to load or load wrongly. Keep groups identical or rebuild the optimiser deliberately.
  • Partial checkpoints. A crash during save leaves an incomplete directory. Write to a temporary path, then mark complete; never overwrite the last good checkpoint.
  • Tokenizer drift. Retokenising new data with a different tokenizer version changes token IDs. Pin the tokenizer file by hash.

Operational guidance

Keep a lineage record for every checkpoint: parent checkpoint, data manifest with hashes, schedule, replay ratio and evaluation results. Without it you cannot answer which data a deployed model saw. Keep at least the last good checkpoint of every increment, because a bad increment is rolled back by restarting from its parent. Budget storage for full optimiser state only on checkpoints you might continue from; older ones can be reduced to exports. Retrain from scratch periodically anyway when the cumulative data shift is large, when architecture or tokenizer changes, or when increments have stacked up beyond what your evaluations can vouch for.

What to do next

  1. Confirm your current checkpoints include FP32 master weights, both Adam moments, scheduler, loader position and per-rank RNG; if not, fix saving before you need it.
  2. Switch to PyTorch DCP with get_state_dict and set_state_dict so checkpoints load on a different GPU count.
  3. Write separate resume and continue code paths, with the continue path building a fresh re-warm and re-decay schedule.
  4. Build a mixed dataset of new data plus replayed old data, deduplicated against the original corpus, with a pinned tokenizer.
  5. Define three evaluation sets, new domain, old distribution and capability, and an acceptance rule for each increment.
  6. Run a small pilot increment at reduced scale to choose the re-warm peak and replay fraction before the full run.
  7. Record checkpoint lineage and keep the parent of every increment until the next one is accepted.
Key takeaway: Incremental training reuses a checkpoint instead of retraining, at a small fraction of the compute, if it restores full optimiser state, re-warms and re-decays the learning rate, replays some old data and evaluates on the old distribution as well as the new. Save sharded checkpoints that reshard on load, keep resume and continue code paths separate, and use adapters when the increment is a task rather than a corpus.