Skip to content

v0.8.0: On-Demand Full-State Checkpointing & OSFT Correctness Fixes

Choose a tag to compare

@RobotSail RobotSail released this 24 Apr 21:27
· 17 commits to main since this release
5d36099

Summary: This release adds signal-driven full-state checkpointing for exact training resumption, fixes a subspace drift bug in OSFT caused by AdamW's element-wise rescaling, and replaces the Gram matrix V projection with a factored form that reduces communication volume by 2–7x.

Highlights

  • 💾 On-Demand Full-State Checkpointing: New --on-demand-checkpointing flag enables signal-driven (SIGTERM, SIGINT, SIGUSR1, etc.) full training state saves (model, optimizer, scheduler, and RNG) for bit-exact resumption via --resume-from-full-state-checkpoint
  • 🔧 OSFT Subspace Leak Fix: Added post-step parameter re-projection to correct for AdamW's weight decay and momentum drifting parameters out of the orthogonal complement. A robustness improvement that makes OSFT correct across a wider range of hyperparameters — no speed penalty, no effect at typical learning rates (1e-5 to 1e-6), but eliminates thousands of orthogonality violations at higher rates
  • ⚡ Factored V Projection: Replaced the Gram matrix V projection $dV \mathrel{-}= dV (V_\text{high}^\top V_\text{high})$, requiring an $M \times M$ all-reduce, with a factored form $dV \mathrel{-}= (dV , V_\text{high}^\top) V_\text{high}$, requiring a $k_\text{high} \times M$ all-gather. 2x fewer bytes for square weights, 7x for down_proj, yielding up to 25% faster OSFT training
  • 🧹 Removed Deprecated osft_memory_efficient_init: The deprecated flag has been fully removed from CLI, TrainingArgs, and api_train.py

New Features

On-Demand Full-State Checkpointing

Add on-demand full-state checkpointing for OSFT training resumption by @RobotSail in #79

  • New FullStateCheckpointer class in full_state_checkpoint.py that installs signal handlers for SIGTERM, SIGINT, SIGUSR1, SIGUSR2, SIGXCPU, SIGHUP, and SIGQUIT, covering Kubernetes, SLURM, PBS, LSF, and interactive sessions
  • Trigger mechanism uses a file on /dev/shm. The signal handler writes the file atomically, workers poll at batch boundaries, and coordinate via all_reduce(MAX) so all ranks agree
  • Zero-gather architecture: model state saved sharded via DCP (each rank saves its own shard), optimizer state per-rank, per-rank RNG snapshots, LR scheduler state, sampler epoch, and training counters
  • RNG state is snapshotted at the start of each training step. If a signal fires mid-gradient-accumulation, partial gradients are discarded and resume replays the full step, guaranteeing bit-identical optimization trajectories
  • Resume path in finalize_model_initialization materializes OSFT parameters via SVD first, then overwrites with checkpoint values via DCP in-place, preserving FSDP2's internal DTensor references and mixed precision tracking
  • Automatic embedding resize on resume when the saved embed_tokens.weight shape differs from the freshly-loaded model (e.g., if the first run resized vocab)
  • Parent process (api_train.py) installs its own signal handler and gives workers up to 300s to save before shutdown
  • Checkpoint checks happen at three points per step: before forward, before backward, and after backward

Usage:

# Training with on-demand checkpointing
torchrun --nproc_per_node=8 -m mini_trainer.train \
    --model-name-or-path meta-llama/Llama-3.1-8B-Instruct \
    --data-path ./data.jsonl \
    --output-dir ./checkpoints \
    --on-demand-checkpointing

# Resuming from a saved checkpoint
torchrun --nproc_per_node=8 -m mini_trainer.train \
    --model-name-or-path meta-llama/Llama-3.1-8B-Instruct \
    --data-path ./data.jsonl \
    --output-dir ./checkpoints \
    --on-demand-checkpointing \
    --resume-from-full-state-checkpoint ./checkpoints/full_state_checkpoints/step_100

Bug Fixes

AdamW Subspace Leak

Fix post-step parameter re-projection to prevent AdamW subspace leak — bug found and fix proposed by @LazarValkov — in #91

This is a correctness fix. The existing pre-step gradient projection is necessary but not sufficient — AdamW's weight decay and element-wise moment rescaling ($\hat{m}_t / \sqrt{\hat{v}_t}$) modify parameters in ways that aren't captured by gradient projection alone, causing trained parameters to drift into the frozen subspace.

  • Added project_parameters() method and project_parameter_to_orthogonal_space() function that projects $U_\text{low}$ and $V_\text{low}$ back into the orthogonal complement after each optimizer step
  • The distributed logic (FSDP2 sharding, all-reduce for U, all-gather for V, optional V_high caching) mirrors the gradient projection exactly
  • Called in the optimizer wrapper immediately after optimizer.step() for OSFT models
  • No training speed penalty or regression on any supported model architecture

Impact: At learning rates used in practice (1e-5 to 1e-6), the fix has no measurable effect on training outcomes — continual learning benchmarks (TRACE, 8 tasks) produce equivalent AA and BWT with and without the fix. At higher learning rates (5e-4), the fix eliminates thousands of orthogonality violations (88% pass rate → 100%). This makes OSFT correct across a wider range of hyperparameters rather than relying on the learning rate being low enough that the drift is negligible.

Validation: Tested on 9 model architectures, orthogonality compliance tests across multiple learning rates, and the TRACE continual learning benchmark.

Multi-Node Logging

Fix LOCAL_RANK for console log gating in #82

  • Changed AsyncStructuredLogger to use LOCAL_RANK instead of global rank for console output gating, so every node's rank-0 process prints training progress in multi-node setups

Performance Improvements

Factored V Projection

Replace Gram matrix V projection with factored form by @stmcgovern in #74

  • Replaced $dV \mathrel{-}= dV (V_\text{high}^\top V_\text{high})$ ($M \times M$ Gram matrix all-reduce) with $dV \mathrel{-}= (dV , V_\text{high}^\top) V_\text{high}$ ($k_\text{high} \times M$ all-gather)
  • Communication savings: $M / k_\text{high}$ fewer bytes. 2x for square weight matrices, up to 7x for down_proj where $k_\text{high} = \min(N, M) \times (1 - \text{URR})$
  • Optional V_high caching via OSFT_CACHE_V=1 environment variable avoids repeating the all-gather on every step (adds ~5.1 GB per rank for Llama-8B; not recommended for 70B+)
  • Handles uneven shards when k_high is not evenly divisible by world_size via padding/slicing

Benchmarks (Granite 3.1 8B, OSFT):

Mode Avg Tokens/s Avg Batch Time (s) Speedup
Gram (baseline) 1,728 1.711 1.00x
Factored 1,936 1.525 1.12x
Factored + Cache 2,160 1.380 1.25x

Chores

  • Relaxed numba pin from >=0.62.0 to >=0.61.2 for compatibility with vllm ≤0.15.1 by @Maxusmusti in #93
  • Removed deprecated osft_memory_efficient_init parameter from CLI, TrainingArgs, and api_train.py in #81
  • Added CLAUDE.md with project context for AI-assisted development by @lukeinglis in #92
  • Added Dependabot config for weekly GitHub Actions auto-updates in #84
  • Added Nemotron compatibility shim for is_flash_attn_greater_or_equal_2_10 renamed in transformers 5.x

Testing

  • New tests/test_full_state_checkpoint.py with 17 unit tests covering signal handling, trigger coordination, metadata save/load round-tripping, and RNG state restore
  • New tests/test_osft.py with extensive OSFT unit tests including SVD decomposition, gradient projection, parameter projection, weight reconstruction, 200-step AdamW stress test, and batched vs. unbatched equivalence
  • New benchmarks/bench_v_proj.py for multi-GPU comparison of Gram vs. factored vs. cached V projection
  • New regression tests for checkpoint fidelity (bit-identical trajectories), multi-node fidelity, signal-driven checkpoint, and all-model matrix testing
  • Improved timing stability in test_api_train.py for CI environments

Upgrade Notes

  • Breaking: The osft_memory_efficient_init CLI flag and TrainingArgs field have been removed. If your training scripts still pass this flag, remove it. It has had no effect since v0.4.0
  • --on-demand-checkpointing defaults to False; opt in explicitly
  • numba minimum version lowered to >=0.61.2 for broader compatibility

Contributors

Installation

Through Pip:

uv pip install rhai-innovation-mini-trainer && uv pip install rhai-innovation-mini-trainer[cuda] --no-build-isolation

Locally:

uv pip install . && uv pip install .[cuda] --no-build-isolation

Full Changelog: v0.7.2...v0.8.0