Repository navigation
Conversation
Keep the four production fixes from vision/SP validation together:
- Use one combined image/video vision-projector call per pack, including
the dummy text-only call. Derive merged feature counts from the raw patch
streams before SP padding while preserving the current gather, trimming,
modality slicing, and placeholder-check interface.
- Keep the non-persistent visual RoPE buffer real on CPU even during meta
construction, explicitly name its FP32 dtype, and document why checkpoint
loading cannot restore it after materialization.
- Replicate ignored FP32 parameters across the full world mesh rather than
only the EP submesh, so gradients, global norms, and clipping agree across
FSDP rows. Preserve existing full-world and sharded placements.
- Gather MTP raw_inputs_embeds with autograd and reduce-scatter SUM in
backward, preserving future-token gradients across SP shard boundaries.
The base already includes the core single-call vision behavior and CPU
RoPE initialization. Those files align the validated patch-count handling
and initialization rationale; replica synchronization and MTP autograd are
the missing behaviors fixed here. No regression test files are included.
Validation:
- Ruff lint/format, Python syntax, and git diff --check passed for all four
production files.
- Ten targeted CPU checks passed for compose call counts, unchanged
embeddings, mocked SP2 padding/modality slices, malformed placeholders,
and real FP32 CPU RoPE buffers under meta construction.
- Prior GPU validation passed the FSDP/EP and MTP gather regressions;
SP1/2/4/8 each completed 20 steps with finite gradients and zero
cross-rank gradient-norm spread. Strict cross-SP equivalence remains
unresolved. GPU tests were not rerun on this branch.
Wuyh11
force-pushed
the
fix/glm53-fsdp-sp-gradients
branch
from
October 9, 2026 11:32
11b18eb to
747b5b5
Compare
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.
What
Keep the four GLM-5.3-Flash training issues identified in vision/SP validation together in this PR, targeting
feat/glm53flash-f6-text-moe(#2111). The production diff coversmodeling_glm53.py,modeling_vision.py,base.py, andsequence_context.py; no regression test files are included.The target branch already contains the core single-call vision behavior and real CPU RoPE initialization. The two vision-file changes align the earlier validated patch-count handling and make the FP32 initialization explicit while preserving the current gather/splice interface. FP32 replica synchronization and MTP embedding autograd are the two missing behaviors fixed by this PR.
1. Keep mixed image/video packs on one vision/projector call
Previously, a rank with both image and video samples entered the vision tower twice, while ranks with one visual modality entered it once and proceeded to the text MoE. The mismatched FSDP vision all-gather and MoE all-to-all order caused the mixed-media training run to hang before any optimizer update.
Keep image/video patches and expanded grids concatenated for one vision/projector call per pack; grid rows remain separate attention segments. Derive each modality's merged feature count from its raw patch stream before SP padding, matching the original validated patch.
_splice()gathers and trims the combined feature stream before taking each modality's rank-local slice and checking placeholder counts. Pure-text packs retain their dummy visual call. The existing_splice()signature and single mask/feature gather are preserved.The single-call fix is already present in the base; the new diff aligns the feature-count calculation and removes the extra grid-reduction
.item()calls.2. Keep visual RoPE frequencies real and FP32 during meta construction
Visual
inv_freqis a non-persistent buffer, so checkpoint loading cannot restore it. Constructing it on meta left uninitialized frequencies after materialization, producing non-finite rotary values, NaN loss/gradients, and skipped optimizer updates.Keep the analytical frequency construction explicitly on CPU and spell its dtype as
torch.float32, with the original patch's initialization rationale documented beside it. The base already has CPU initialization; this change clarifies the dtype and buffer lifecycle rather than claiming another new NaN fix.3. Synchronize ignored FP32 parameters across FSDP rows
BaseModel._fully_shard()previously kept an existing EP-submeshReplicateplacement when adding a FP32 parameter to FSDP'signored_params. EP-only replication leaves independent copies on other FSDP rows, so gradient reduction missed those rows and ranks used different global norms and clipping factors.The production check identified 28 affected parameters. With four local gradients
[1, 2, 3, 4]on FSDP2 x EP2, the old path produced group means1.5and3.5instead of the required global mean2.5.Fix: redistribute submesh replicas onto the full world mesh before ignoring them in FSDP. Preserve existing full-world replicas and genuinely sharded placements; detach the local DTensor value before redistribution so the new parameter is a leaf tensor.
4. Preserve MTP future-embedding gradients under SP
MTP rolls
SequenceContext.raw_inputs_embedsto construct future-token inputs. SP1 returned the original embeddings with their graph, but SP>1 used ordinarydist.all_gather, whose receive tensors were detached from the local embeddings. The MTP embedding-input branch silently lost its gradients, including dependencies crossing shard boundaries. The main hidden-state branch still backpropagated, so a close loss or global norm did not detect the error.Fix: use
all_gather_tensor_autogradfor embeddings, whose backward reduce-scatters with SUM to return gradients to their owners. Keep the cached tensor attached to its graph. Integer token-ID gathering is unchanged.Tests
git diff --checkpass.pre_commitis not installed. GPU checks were not rerun on the latest branch.Verification (8xH200, prior fix validation)
The following results are from September 28-30 validation of the original four fixes on an earlier GLM53 training snapshot, not a fresh GPU run of this PR head.
1.5/3.5) and passed afterward (2.5on every rank), including matching norms and clipping results.Gradient comparisons cover all 347 vision/projector parameters before clipping. Strict numerical equivalence across SP sizes is still not achieved. Later diagnostics identified sparse dKV backward nondeterminism and projector batch-shape rounding as residual sources; those kernels and the balancing-loss objective are unchanged. No before/after throughput comparison was run.
The separate TileLang selector issue tracked by #2143 must also be fixed on the stack before a new default-backend GPU end-to-end validation; it is outside these four vision/parallel-training issues.