fix(train): serialize GaLore projector state in checkpoints - #10161
Open
MrCapricornLiu wants to merge 1 commit into
Open
MrCapricornLiu wants to merge 1 commit into
MrCapricornLiu wants to merge 1 commit into
Conversation
tastelikefeet
approved these changes
Sep 22, 2026
Collaborator
|
lint the code by: please |
This branch has not been deployed
This file contains hidden or bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters.
Learn more about bidirectional Unicode characters
Sign up for free
to join this conversation on GitHub.
Already have an account?
Sign in to comment
Add this suggestion to a batch that can be applied as a single commit.This suggestion is invalid because no changes were made to the code.Suggestions cannot be applied while the pull request is closed.Suggestions cannot be applied while viewing a subset of changes.Only one suggestion per line can be applied in a batch.Add this suggestion to a batch that can be applied as a single commit.Applying suggestions on deleted lines is not supported.You must change the existing code in this line in order to create a valid suggestion.Outdated suggestions cannot be applied.This suggestion has been applied or marked resolved.Suggestions cannot be applied from pending reviews.Suggestions cannot be applied on multi-line comments.Suggestions cannot be applied while the pull request is queued to merge.Suggestion cannot be applied right now. Please check back later.
GaLore stores a
GaLoreProjectorinstance in optimizer state. Consequently, a normal optimizer checkpoint cannot be restored by Trainer'storch.load(..., weights_only=True): it fails with an unsupported-global error before training resumes.Serialize the projector as a dictionary of configuration values and tensors, then reconstruct it when loading optimizer state. The live optimizer keeps its projector object. Move cached projection bases to the gradient's device/dtype before reuse, so CPU-mapped checkpoints can resume on CUDA without waiting for the next SVD refresh. This applies to the bundled AdamW, Adafactor and AdamW8bit implementations.
Validation:
git diff --checkpass. Base-revision whole-repository lint has an unrelated formatting failure intests/megatron/test_infonce_ddp_e2e.py.This makes newly written checkpoints compatible with safe loading. It does not change safe-loading policy for older pickled-object checkpoints, or address the separate per-parameter optimizer wrapper's checkpoint handling. Validation used FP32 parameters; distributed/sharded checkpoints and q-galore were not tested.