Skip to content

[Fix] Align GLM-5.3-Flash vision training and parallel gradients - #2144

Open
Wuyh11 wants to merge 1 commit into
InternLM:feat/glm53flash-f6-text-moefrom
Wuyh11:fix/glm53-fsdp-sp-gradients
Open

Wuyh11 wants to merge 1 commit into
InternLM:feat/glm53flash-f6-text-moefrom
Wuyh11:fix/glm53-fsdp-sp-gradients

Conversation

@Wuyh11

@Wuyh11 Wuyh11 commented Oct 9, 2026 •

Copy link
Copy Markdown

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 covers modeling_glm53.py, modeling_vision.py, base.py, and sequence_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_freq is 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-submesh Replicate placement when adding a FP32 parameter to FSDP's ignored_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 means 1.5 and 3.5 instead of the required global mean 2.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_embeds to construct future-token inputs. SP1 returned the original embeddings with their graph, but SP>1 used ordinary dist.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_autograd for 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

  • Current four-file diff: Ruff lint/format, Python syntax, and git diff --check pass.
  • Ten targeted CPU checks pass using the actual compose/RoPE methods: text, image, video and mixed packs; both SP2 shards with a modality boundary and padding; placeholder-count rejection; and real finite FP32 CPU frequencies under a meta construction context. Embeddings are unchanged from the base in the valid compose cases. SP collectives and the vision encoder were mocked, so these checks do not replace GPU validation.
  • Prior local regression coverage includes the four-rank FSDP2 x EP2 gradient/norm/clipping check and two GPU MTP gather checks. These regression files are not committed.
  • Full pre-commit could not run because pre_commit is not installed. GPU checks were not rerun on the latest branch.
  • Both ReadTheDocs checks currently report failure. The same contexts also fail on the base PR [Feature] Add GLM-5.3-Flash F6 core: text model + MTP + compose model #2111; the cause has not been confirmed from the available build data.

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.

  • Mixed-media execution progressed past the original communication hang. After real FP32 RoPE initialization, loss/gradients were finite and optimizer updates occurred: SP1 completed five updates, SP2 two, and frozen-vision SP1 two. The frozen vision parameters remained unchanged.
  • The four-rank FSDP/EP regression failed before the fix (1.5 / 3.5) and passed afterward (2.5 on every rank), including matching norms and clipping results.
  • Both MTP GPU regressions passed; cross-shard packed-roll input gradients matched the full-sequence reference elementwise.
  • With BF16, EP4, global batch 8, 16K packs, MTP coefficient 0.1, and balance alpha 0.001, SP1/2/4/8 each completed 20 optimizer steps. Every rank completed its run; gradients were finite and present, optimizer state advanced, and all 80 training steps had zero cross-rank gradient-norm spread.
SP Completed optimizer steps Step-20 loss First-step visual gradient relative L2 vs SP1
1 20 10.071869 Reference
2 20 10.069990 3.468954%
4 20 10.079679 3.283411%
8 20 10.100826 1.267570%

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.

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
Wuyh11 force-pushed the fix/glm53-fsdp-sp-gradients branch from 11b18eb to 747b5b5 Compare October 9, 2026 11:32
@Wuyh11 Wuyh11 changed the title [Fix] Preserve GLM-5.3-Flash gradients across FSDP and SP [Fix] Align GLM-5.3-Flash vision training and parallel gradients Oct 9, 2026

This branch has not been deployed

No deployments
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

None yet

Projects

None yet

Development

Successfully merging this pull request may close these issues.

1 participant