Repository navigation
v0.8.0: On-Demand Full-State Checkpointing & OSFT Correctness Fixes
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-checkpointingflag 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, andapi_train.py
New Features
On-Demand Full-State Checkpointing
Add on-demand full-state checkpointing for OSFT training resumption by @RobotSail in #79
- New
FullStateCheckpointerclass infull_state_checkpoint.pythat 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 viaall_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_initializationmaterializes 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.weightshape 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_100Bug 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 (
- Added
project_parameters()method andproject_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
AsyncStructuredLoggerto useLOCAL_RANKinstead 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=1environment 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.0to>=0.61.2for compatibility with vllm ≤0.15.1 by @Maxusmusti in #93 - Removed deprecated
osft_memory_efficient_initparameter from CLI,TrainingArgs, andapi_train.pyin #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_10renamed in transformers 5.x
Testing
- New
tests/test_full_state_checkpoint.pywith 17 unit tests covering signal handling, trigger coordination, metadata save/load round-tripping, and RNG state restore - New
tests/test_osft.pywith 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.pyfor 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.pyfor CI environments
Upgrade Notes
- Breaking: The
osft_memory_efficient_initCLI flag andTrainingArgsfield have been removed. If your training scripts still pass this flag, remove it. It has had no effect since v0.4.0 --on-demand-checkpointingdefaults toFalse; opt in explicitly- numba minimum version lowered to
>=0.61.2for broader compatibility
Contributors
Installation
Through Pip:
uv pip install rhai-innovation-mini-trainer && uv pip install rhai-innovation-mini-trainer[cuda] --no-build-isolationLocally:
uv pip install . && uv pip install .[cuda] --no-build-isolationFull Changelog: v0.7.2...v0.8.0