Skip to content

[Feature] Add GLM-5.3-Flash F6 core: text model + MTP + compose model - #2111

Open
jayhenry wants to merge 10 commits into
feat/glm53flash-f2-vision-towerfrom
feat/glm53flash-f6-text-moe
Open

jayhenry wants to merge 10 commits into
feat/glm53flash-f2-vision-towerfrom
feat/glm53flash-f6-text-moe

Conversation

@jayhenry

@jayhenry jayhenry commented Sep 23, 2026 •

Copy link
Copy Markdown
Collaborator

Stack (bottom to top):

  1. [Feature] Add GLM-5.3-Flash F0: 25B cropped reference checkpoint builder #2105 feat/glm53flash-materialize-full-f0 → main
  2. [Feature] Add GLM-5.3-Flash F3: Kimi Delta Attention (KDA) #2106 feat/glm53flash-f3-kda → feat/glm53flash-materialize-full-f0
  3. [Feature] Add GLM-5.3-Flash F4: mHC four-stream residual #2107 feat/glm53flash-f4-mhc → feat/glm53flash-f3-kda
  4. [Feature] Add GLM-5.3-Flash F5: NoPE DSA + KPool indexer + clamped SwiGLU #2108 feat/glm53flash-f5-nope-dsa → feat/glm53flash-f4-mhc
  5. [Feature] Add GLM-5.3-Flash F1: VL data preprocessing pipeline #2109 feat/glm53flash-f1-vl-data → feat/glm53flash-f5-nope-dsa
  6. [Feature] Add GLM-5.3-Flash F2: vision tower + projector (eager) #2110 feat/glm53flash-f2-vision-tower → feat/glm53flash-f1-vl-data
  7. [Feature] Add GLM-5.3-Flash F6 core: text model + MTP + compose model #2111 feat/glm53flash-f6-text-moe → feat/glm53flash-f2-vision-tower ← you are here

Base is #2110's branch (layer 6). Review only this PR's own diff.


Summary

Stack layer 7/7 (top) of GLM-5.3-Flash support (base: layer 6, F2 vision tower). Covers the F6 text model + MTP + compose model, and the multimodal enablement that makes VL SFT actually runnable: sequence-parallel splice, and the VL training config/launcher.

Implements the core of the F6 milestone: the 45-layer KDA/NoPE-DSA text model with mHC four-stream residual, MTP, and the compose model wiring vision_tower + multi_modal_projector + language_model with image/video splice.

  • xtuner/v1/model/moe/glm53/glm53.py: Glm53TextMoEConfig/Glm53TextMoE. Layer schedule ([KDA,KDA,KDA,DSA]x11 + KDA) read from the real checkpoint's text_config.layer_types. mHC's four-stream residual is expanded/collapsed once at the decoder-stack boundary, not per layer, so MoE._forward's aux-loss/router bookkeeping is reused unmodified. MTP reuses the base MTPLayer/MTPBlock unmodified — the design doc's originally-planned mtp.py deliverable turned out unnecessary.
  • xtuner/v1/model/compose/glm53/modeling_glm53.py + Glm53BaseConfig: Glm53ForConditionalGeneration. Image/video splice uses the global mm_token_type_ids, never input_ids==video_token_id. A placeholder-count mismatch raises immediately (design doc §16.2) — unlike Qwen3-VL's compose path, this never silently continues training on a corrupted splice.
  • Registers the "glm5.3" chat template / tokenize-fn branch, get_model_config_from_hf dispatch, and the SFT training config/launch script (examples/v1/config/sft_glm53.py, sft_glm53_tiny.sh) for the 8-GPU end-to-end smoke suite.

Root-caused four bugs found by end-to-end gradient-flow and 8-GPU training smoke tests (minimal repros, not assumed):

  1. FLA's chunk_kda Triton kernel silently drops backward gradients when head_dim < 16 (a tl.dot minimum-K constraint) — not a production issue since real GLM-5.3-Flash uses head_dim=128.
  2. DSA indexer showing no gradient with freeze_dsa_indexer=False traced to an existing, intentional torch.no_grad() wrap (not a bug).
  3. torch.compile traced KDA's Triton-kernel-heavy forward under a strict boundary and hit disabled-function errors; fixed with GLM53_MOE_{NON_EP,EP}_COMPILE_CFG. Surfaced a separate real upstream limitation: FLA's prepare_chunk_indices/prepare_lens do a .tolist()-driven Python loop over cu_seqlens, incompatible with dynamo's dynamic-shape tracing (_mark_dynamic on the packed boundaries). Initially worked around by wrapping FLA's entry points in torch._dynamo.disable(), which made compile run but cost ~2.4x throughput; now fixed properly by owning the kernel behind torch.library.custom_op — see torch.compile throughput below.
  4. Under ep_size>1, KDA's o_norm (FLA's FusedRMSNormGated) crashed with a Triton illegal-memory-access — its inherited forward reads self.weight directly, which stays a DTensor under EP. Fixed with a minimal forward override that unshards weight/bias first.

Known gaps (recorded explicitly, not silently skipped): the compose layer has no dedicated multi-GPU FSDP parity test beyond the smoke suite and the 2-GPU SP splice test below.

Test Plan

  • tests/model/test_glm53_text_moe.py::TestGlm53TextMoEAccuracy::test_fsdp_accuracy: XTuner forward loss vs real transformers.Glm5NextForConditionalGeneration on the F0 25B cropped checkpoint, for (dispatcher, ep_size) in {(None, 1), ("all2all", 4), ("all2all", 8)} — all pass. The ep_size=4/8 cases double as the regression test for bug 4 above.
  • sft_glm53_tiny.sh: single-node 8-GPU end-to-end SFT, real F0 checkpoint, PACK_MAX_LENGTH=16384, TOTAL_STEP=20. Default profile plus five combos (SP_SIZE=2 EP_SIZE=4; EP_SIZE=8 SP_SIZE=1; XTUNER_ACTIVATION_OFFLOAD=0; FP8=1; MODEL_COMPILE=1) all complete 20 steps with loss monotonically decreasing, no NaN/OOM. See the table below.

Text end-to-end regression (sft_glm53_tiny.sh, single-node 8-GPU, real F0 25B cropped checkpoint)

PACK_MAX_LENGTH=16384, TOTAL_STEP=20. All six profiles complete 20 steps, loss monotonically decreasing, no NaN/OOM. Numbers are rank0; mem is max_memory / reserved_memory (GB); tgs/seqlen_tgs/exp_tgs are step-20 throughput (tokens/gpu/s); time is step1's own duration (includes one-time warmup/compile) plus the step1→step20 wall-clock gap:

profile step1 loss step20 loss grad_norm@20 mem@20 (GB) tgs@20 seqlen_tgs@20 exp_tgs@20 time
default (EP4 SP1 offload=1) 11.30479 10.31409 48.41 86.22 / 106.10 12203.1 12224.0 4212.1 76s
SP2 EP4 11.31528 10.32397 48.40 75.15 / 96.00 9631.7 9631.7 2617.5 61s
EP8 SP1 11.30479 10.31380 48.18 100.44 / 112.75 11865.9 11886.2 4692.9 68s
offload=0 (EP4 SP1) 11.30479 10.31401 48.41 85.58 / 105.69 12315.3 12336.4 4944.9 64s
FP8=1 (EP4 SP1 offload=1) 11.30908 10.18042 45.88 84.73 / 104.35 6159.6 6170.2 1871.5 174s
MODEL_COMPILE=1 (EP4 SP1 offload=1) 11.30469 10.31426 48.39 82.11 / 103.28 13381.1 13400.0 2270.0 86s
  • SP2's per-GPU effective seqlen is halved (8192 vs 16384), so tgs is lower and wall-clock per step is also shorter than the default despite the added SP communication.
  • offload=0 is faster with higher tgs than the default (offload=1) — expected, since activation CPU-GPU transfer is skipped.
  • EP8 tracks the default (EP4) on loss/tgs; higher mem is because dp_size = world_size/ep_size drops from 2 to 1, so FSDP shards non-expert params over fewer dp ranks.
  • default/EP8/offload=0 step1/step20 loss match to 1e-4, as expected — these knobs change parallelism/memory strategy, not the math.
  • FP8=1's step20 loss differs from the bf16 baseline (10.18 vs 10.31) because FP8 genuinely changes the numerics on the (unquantized-kv_b_proj) absorbed-MLA path, not a bug; its tgs is roughly half the baseline's at this 25B/20-step smoke scale, consistent with per-tensor quantize/dequantize overhead not yet amortized.
  • MODEL_COMPILE=1 is now the fastest profile: steady-state 1.21s/step vs eager's 1.34s (tgs 13381 vs 12203, ~+10%), at lower peak memory. step20 loss (10.31426) matches the eager baseline (10.31409) to 1.8e-4, confirming compile does not change the math.

Multimodal enablement

Sequence parallel in the compose model

The compose model asserted sequence_parallel_mesh was None or size 1, so VL training could not use SP at all — the last piece of the Vision SP gap, now that the tower (layer 6) shards patches merge-aligned.

The splice is the only part that needs to know about SP. input_ids and mm_token_type_ids arrive already sharded (they are split together), so the local mask is what indexes the local embeddings; what the local mask cannot say is which global features belong to this rank. Image and video patches are concatenated into one tower call. _splice gathers the mask once, _gather_visual_features reassembles the globally-ordered features (trimming the tower's merge-alignment padding), and each modality is sliced at the offset of same-modality placeholders on preceding ranks. The gather is the autograd-aware one, so its backward reduce-scatters and hands each rank the gradient of exactly the shard it produced.

The placeholder↔feature count check now compares global totals, so a corrupted sample raises identically with and without SP. The pure-text dummy visual forward stays replicated: its gradient contribution is exactly zero, so sharding it would only add collectives.

VL training entry point

Glm53BaseConfig.from_hf raised NotImplementedError, so the VL model could not be built from a checkpoint at all. It now reads the checkpoint's single vision_config and fans it out to both XTuner modules — tower and projector are one HF module that §5.2 splits in two — and delegates the text half to Glm53TextMoEConfig.from_hf.

examples/v1/config/sft_glm53_vl.py + sft_glm53_vl_tiny.sh mirror the text pair knob for knob, so both modalities run the same acceptance matrix. Media comes from the ci_vl corpus (image + video samples). FP8 stays on the language tower only, matching how the published checkpoint stores the vision half in bf16.

hf_config intentionally remains None: _write_hf_index_and_config then copies the source checkpoint's config/tokenizer into the export, which is self-consistent for a fine-tune that does not change the architecture.

VL end-to-end regression (sft_glm53_vl_tiny.sh, single-node 8-GPU, real F0 25B cropped checkpoint)

PACK_MAX_LENGTH=16384, SAMPLE_MAX_LENGTH=16384, TOTAL_STEP=20, image + video data. All six profiles complete 20 steps, loss decreasing, no NaN/OOM. Numbers are rank0; mem is max_memory / reserved_memory; time is step 1's duration plus the step1→step20 wall-clock gap. This replaces the earlier table: that one trained no real video (see #2109).

profile step1 loss step20 loss grad_norm@20 mem@20 (GB) tgs@20 seqlen_tgs@20 exp_tgs@20 time img_tokens
default (EP4 SP1 offload=1) 12.48675251 10.78624153 25.85 90.59 / 111.43 4465.1 5084.9 2740.3 67s 56424
SP2 EP4 12.21398830 11.29156303 40.15 77.09 / 97.02 3674.0 3674.0 1926.0 44s 31884
EP8 SP1 12.48675251 10.78038311 47.16 103.60 / 121.47 4468.9 5089.2 2563.4 71s 56424
offload=0 (EP4 SP1) 12.48675251 10.81396198 24.98 90.25 / 111.00 4488.9 5112.0 2701.5 69s 56424
FP8=1 12.49881649 10.27666283 23.44 89.44 / 110.15 3972.2 4523.6 2886.5 70s 56424
MODEL_COMPILE=1 12.48672199 10.78443432 25.04 86.27 / 109.73 5727.1 6522.1 2366.0 62s 56424
  • tgs counts real tokens (mask.sum()), seqlen_tgs counts every slot including padding (mask.numel()). The previous table had seqlen_tgs ≈ 2× tgs because packs were half empty: 17% of samples were cached at 4096 tokens and trained as 33-token fakes, so no video sample ever ran ([Feature] Add GLM-5.3-Flash F1: VL data preprocessing pipeline #2109). With videos actually in the pack, rank0 fill on the non-SP profiles is 87–99% (mean 93%) and step-20 seqlen_tgs / tgs is 1.14. The remainder is the tail of a 16384 pack that the next ~14k-token video does not fit into. SP2 packs at 8192 and fills exactly, so the two rates match.
  • default / EP8 / offload=0 step-1 loss is bit-identical (12.48675251).
  • SP2 is not row-comparable. GLOBAL_BATCH_SIZE = world/sp and the sequence is split, so a step sees about half the tokens (img_tokens 31884 vs 56424) and step-20 loss stays higher (11.29 vs 10.79). Vision SP correctness is the 2-GPU bitwise test, not this smoke.
  • The first baseline after the data fix hung on step 4. A pack that holds both an image sample and a video sample called the vision tower once per modality, while a rank with one modality called it once. The tower's FSDP mesh is the world group (8) and the language tower's is the dp group (2), so rank 0 was blocked in the language all-gather (NumelOut/NumelIn = 2) while the other ranks were still in a second vision all-gather (ratio 8) until the NCCL watchdog fired at 30 min. Image and video patches are now one tower call per pack. Each grid_thw row stays its own cu_seqlens segment, and the mixed-pack logits match splicing each modality from its own call, bit for bit.
  • MODEL_COMPILE=1 is the fastest row here (tgs 5727 vs 4465) and the lowest peak memory. step-1 loss matches eager to 3e-5.
  • FP8=1's step-20 loss is lower (10.28 vs 10.79), the same direction as text-side FP8: the absorbed-MLA path's numerics change.
  • EP8 peaks highest because dp_size = world_size/ep_size drops from 2 to 1.

torch.compile throughput

Review caught that the compile row was originally slower than eager (tgs 5168 vs 12203) — backwards from what compile is for. The number was real (identical config, and identical per-step text_tokens between the two runs), but the explanation first written here ("few steps to amortize warmup") was wrong: eager settles at 1.34s/step by step 8 and holds within 0.5%, while compile never converged, bouncing 2.6–4.3s through step 20.

Diagnosis by measurement: TORCH_LOGS=recompiles ruled out runaway recompilation (200 recompiles, but bounded — they stop by step 2, dynamo cache versions cap at /3, and the guard failure is requires_grad mismatch, the expected one-time doubling from activation-checkpoint recompute). TORCH_LOGS=graph_breaks then put 32 breaks exactly on the torch._dynamo.disabled FLA entry points. GLM-5.3-Flash is KDA-dominated (34 of 45 layers), so every KDA layer broke the enclosing compiled region, leaving fragments whose guard and re-entry cost exceeded the fusion benefit.

Fixed in the F3/KDA commit by taking the route XTuner's own GatedDeltaNet already takes (and the reason Qwen3.5 can run GatedDeltaNet.forward at fullgraph=True): xtuner/v1/ops/kda calls FLA's chunk_kda_fwd/chunk_kda_bwd behind torch.library.custom_op + register_fake, so dynamo traces through them opaquely and prepare_chunk_indices stays hidden without a break.

text model steady per-step tgs@20 mem@20 step20 loss
eager 1.34s 12203.1 86.22 10.31408882
compile, before 2.6–4.3s (never converges) 5168.4 82.15 10.31412411
compile, after 1.21s 13381.1 82.11 10.31426430

Verified bitwise-identical to fla.ops.kda.chunk_kda — forward and all five gradients at max|diff| = 0. That check earned its keep: the first attempt fed the original q/k to l2norm_bwd where FLA feeds the L2-normed ones, corrupting only the q/k gradients (~1.0–1.5 max diff) while the forward stayed bitwise correct.

The 4353 → 4090 VL comparison was on the pre-fix corpus, which never trained a real video. The VL table above replaces it: with videos in the pack, compile leads eager (tgs 5727 vs 4465) and uses less memory (86.3 vs 90.6 GB).

fused_kda_gate and the short convolution are ported the same way (kda_gate_fwd/bwd, causal_conv1d_fwd/bwd behind custom_op), bitwise with FLA on the forward and every gradient. Breaks fell from 30 to 10 and throughput did not move (tgs 13376 vs 13381, steady 1.21s/step). The regression was the chunk_kda break; break count is the wrong metric, location is the right one. The KPool num_pools site stays unported: it is a real data-dependent shape, and pushing ctx.new_dynamic_size() into the indexer is likely a net loss.

VL unit coverage

TestGlm53ComposeSequenceParallel::test_image_splice_under_sp_matches_non_sp (2 GPU): two placeholder spans placed so one lands on each rank's shard — forcing the cross-rank feature slicing — and each rank's logits match the corresponding slice of the non-SP run. tests/model/test_glm53_compose.py is 9 passed on 2 GPUs, including test_every_pack_calls_the_vision_tower_exactly_once (text, image, image+video) and the mixed-pack bitwise splice check.

jayhenry and others added 8 commits September 24, 2026 20:22
Implements the core of the F6 milestone from doc/xtuner_glm5p3flash_design.md:
the 45-layer KDA/NoPE-DSA text model with mHC four-stream residual, MTP,
and the compose model wiring vision_tower + multi_modal_projector +
language_model with image/video splice.

- xtuner/v1/model/moe/glm53/glm53.py: Glm53TextMoEConfig/Glm53TextMoE.
  Layer schedule ([KDA,KDA,KDA,DSA]x11 + KDA) read from the real
  checkpoint's text_config.layer_types via build_layers() dispatching
  per-layer on config.attention (NoPEDSAMLAConfig) vs
  config.linear_attention (KDAConfig). mHC's four-stream residual is
  expanded/collapsed once at the _decoder_stack/_micro_batch_decoder_stack
  boundary (not per layer -- each Glm53*DecoderLayer, already built in
  F4, handles its own hc_pre/hc_post internally), so MoE._forward's
  aux-loss/router bookkeeping is reused completely unmodified. No
  cross-layer dsa_topk_ids IndexShare is needed (every indexer_types
  entry in the real checkpoint is "full": 11 main-stack DSA layers + 1
  MTP, confirmed by inspection) -- MTPLayer/MTPBlock are used
  unmodified via build_mtp_block(mhc_cfg=None), which is also why the
  design doc's listed mtp.py deliverable turned out to be unnecessary:
  the base MTPLayer's enorm/hnorm/eh_proj/final_layernorm param names
  already match checkpoint layers.45 exactly.
- xtuner/v1/model/compose/glm53/modeling_glm53.py +
  Glm53BaseConfig (glm53_config.py): Glm53ForConditionalGeneration.
  Image and video get separate vision-tower forwards; splice uses the
  global mm_token_type_ids (1=image, 2=video), never
  input_ids==video_token_id (F1.b: that token never appears in the
  expanded sequence). A placeholder-count mismatch raises immediately
  -- unlike Qwen3-VL's compose path, this never catches the mismatch
  in a bare except and continues training on a corrupted splice
  (design doc §16.2).
- xtuner/v1/data_proto/templates/__init__.py,
  xtuner/v1/datasets/sft_tokenize_fn/openai.py: registers the
  "glm5.3" chat template / Glm53ChatMessages tokenize-fn branch so a
  generic OpenaiTokenizeFunctionConfig(chat_template="glm5.3") works.
- xtuner/v1/model/__init__.py: get_model_config_from_hf dispatches
  glm5_next's model_type to Glm53TextMoEConfig.from_hf (text-only SFT
  path; the VL compose config covers the vision half separately).
- examples/v1/config/sft_glm53.py, sft_glm53_tiny.sh: SFT training
  config and launch script for the 8-GPU end-to-end smoke suite,
  adapted from GLM-5.2's equivalents (glm5.3 chat template,
  flash_mla_cudnn default sparse_mla_backend, no
  resolve_indexer_topk_query_chunk_size -- NoPEDSAMLAConfig has no
  "deep_gemm_fp8" indexer backend).

Root-caused four bugs found by end-to-end gradient-flow and 8-GPU
training smoke tests, via minimal repros rather than assuming any was
real:
(1) KDA's core params showed no gradient in a small synthetic test --
isolated to FLA's chunk_kda Triton kernel silently dropping backward
gradients when head_dim < 16 (a tl.dot minimum-K constraint), not a
production issue since real GLM-5.3-Flash uses head_dim=128.
(2) The DSA indexer showed no gradient even with
freeze_dsa_indexer=False -- traced to nope_dsa_mla.py's indexer call
being unconditionally wrapped in torch.no_grad() (existing F5 code),
the standard DSA-indexer training convention, not a bug.
(3) Glm53TextMoE had no default_compile_cfg override, so
torch.compile traced KDA's Triton-kernel-heavy forward under a strict
boundary and hit `torch.compiler.disable()`d-function errors; fixed
by adding GLM53_MOE_NON_EP_COMPILE_CFG/GLM53_MOE_EP_COMPILE_CFG
(mirroring GLM-5.2's pattern). This then surfaced a separate, real
upstream limitation: FLA's prepare_chunk_indices/prepare_lens
(fla/ops/utils/index.py) does a .tolist()-driven Python loop over
cu_seqlens, incompatible with dynamo's dynamic-shape tracing --
confirmed via isolated repro, not fixed (MODEL_COMPILE=0 used for the
8-GPU smoke suite; compile isn't in the acceptance criteria's required
coverage).
(4) Under ep_size>1, KDA's o_norm (FLA's FusedRMSNormGated) crashed
with a Triton "illegal memory access" reading address 0x0 --
CUDA_LAUNCH_BLOCKING=1 + compute-sanitizer --tool memcheck localized
it to a null data_ptr(): unlike A_log/dt_bias/the conv weight (all
explicitly unsharded via _to_local() before use), FusedRMSNormGated's
inherited forward reads self.weight directly, which stays a DTensor
under EP (ep_size=1's FSDP2 pre-forward hook happened to auto-unshard
it, masking the bug). Fixed with a minimal FusedRMSNormGated.forward
override in kda.py that unshards weight/bias first, mirroring the
existing _to_local() pattern.

Test plan:
- tests/model/test_glm53_text_moe.py, test_glm53_compose.py: unit and
  real-checkpoint weight-coverage tests, GPU.
- tests/model/test_glm53_text_moe.py::TestGlm53TextMoEAccuracy::
  test_fsdp_accuracy (验收 1): XTuner forward loss vs real
  transformers.Glm5NextForConditionalGeneration on the F0 25B cropped
  checkpoint, for (dispatcher, ep_size) in {(None, 1), ("all2all", 4),
  ("all2all", 8)} -- all pass. Installed transformers has no MTP
  forward for this class, so both sides compare only the 5-layer main
  stack (mtp_config=None on the XTuner side); sparse_mla_backend/
  indexer_backend forced to "torch" (eager) since the test sentences
  are shorter than flash_mla_cudnn's 512-token alignment. The
  ep_size=4/8 cases double as the real-checkpoint regression test for
  bug (4) above.
- sft_glm53_tiny.sh (验收 2): single-node 8-GPU end-to-end SFT, real
  F0 checkpoint, PACK_MAX_LENGTH=16384, TOTAL_STEP=20, MODEL_COMPILE=0.
  Default profile (EP_SIZE=4 SP_SIZE=1 XTUNER_ACTIVATION_OFFLOAD=1)
  and the three required combos (SP_SIZE=2 EP_SIZE=4; EP_SIZE=8
  SP_SIZE=1; XTUNER_ACTIVATION_OFFLOAD=0) all complete 20 steps with
  loss monotonically decreasing (~11.3 -> ~10.3), no NaN/OOM.

Known gaps (recorded explicitly, not silently skipped): Vision SP
remains the gap F2 already recorded (not implemented, not just
untested), and the compose splice assumes sequence_parallel_mesh is
None/size 1; FSDP2/compile/FP8 paths have no multi-GPU test in the
compose layer; FP8 (FP8=1) was not included in the 8-GPU smoke combos,
only FP8=0; MODEL_COMPILE=1 end-to-end training does not work, see bug
(3) above (upstream FLA/dynamo limitation, not required by the
acceptance criteria).

Co-Authored-By: Claude Sonnet 5 <noreply@anthropic.com>
…odel

The compose model asserted `sequence_parallel_mesh` was None or size 1, so VL training could
not use SP at all -- the last piece of the Vision SP gap, now that the tower shards patches
merge-aligned.

The splice is the only part that needs to know about SP. `input_ids` and `mm_token_type_ids`
arrive already sharded (they are split together, F1.b), so the local mask is what indexes the
local embeddings; what the local mask cannot say is *which* global features belong to this rank.
`_local_feature_slice` therefore gathers the mask to compute this rank's offset (the number of
same-modality placeholders on all preceding ranks), and `_gather_visual_features` reassembles
the globally-ordered features, trimming the tower's merge-alignment padding. The gather is the
autograd-aware one, so its backward reduce-scatters and hands each rank the gradient of exactly
the shard it produced.

The placeholder<->feature count check now compares *global* totals, so a corrupted sample raises
the same way with and without SP rather than turning into a per-rank count mismatch.

The pure-text dummy visual forward stays replicated: its gradient contribution is exactly zero
(`dummy_feats.sum() * 0.0`), so sharding it would only add collectives.

Test Plan:
  - `TestGlm53ComposeSequenceParallel::test_image_splice_under_sp_matches_non_sp` (2 GPU): two
    placeholder spans placed so one lands on each rank's shard -- forcing the cross-rank feature
    slicing -- and each rank's logits match the corresponding slice of the non-SP run. This
    replaces the guard test that asserted SP was refused.
  - tests/model/test_glm53_compose.py 6/6.

Co-Authored-By: Claude Opus 5 <noreply@anthropic.com>
Multimodal training had no entry point: `Glm53BaseConfig.from_hf` raised `NotImplementedError`,
so the VL model could not be built from a checkpoint at all, and there was no training config or
launch script to pair with the text-only `sft_glm53_tiny.sh`.

`Glm53BaseConfig.from_hf` reads the checkpoint's single `vision_config` and fans it out to both
XTuner modules -- the tower and the projector are one HF module that §5.2 splits in two -- and
delegates the text half to `Glm53TextMoEConfig.from_hf`. `rope_parameters` is left at its default
because the published config.json has no such key and HF's `AutoConfig` fills the same value.

`examples/v1/config/sft_glm53_vl.py` + `sft_glm53_vl_tiny.sh` mirror the text pair knob for knob,
so the two modalities run the same acceptance matrix. Media comes from the `ci_vl` corpus that
zdev/sft_qwen35_mengke.sh uses (image + video samples). FP8 stays on the language tower only,
matching how the published checkpoint stores the vision half in bf16.

Test Plan: single-node 8-GPU run on the F0 25B cropped checkpoint completes with
`PACK_MAX_LENGTH=16384` (~31k image patches per step), loss decreasing, 85.85 GB peak. Full
profile matrix recorded in doc/progress.md.

Co-Authored-By: Claude Opus 5 <noreply@anthropic.com>
Adds an H2 for the 2026-09-24 round: the contract-by-contract comparison against Automodel's
GLM-5.3-Flash SFT implementation (text side matched throughout, nothing to change), the five real
bugs that only surfaced by running VL training, the Vision SP design as built, the ViT activation
recompute that was the actual OOM cause, and the six-profile VL acceptance matrix.

Also records why `hf_config` staying None is a decision rather than a gap: with it None,
`_write_hf_index_and_config` copies the source checkpoint's config, which is self-consistent for
a fine-tune that does not change the architecture.

Co-Authored-By: Claude Opus 5 <noreply@anthropic.com>
`MODEL_COMPILE=1` made the text model ~2.4x *slower* per step than eager -- the opposite of what
compile is for. Measured on identical data and config (only MODEL_COMPILE differing): eager
settles at 1.34s/step by step 8 and holds within 0.5%, while compile never converged, bouncing
2.6-4.3s through step 20.

Root cause, established by measurement rather than inspection: `TORCH_LOGS=recompiles` ruled out
runaway recompilation (200 recompiles, but they stop by step 2 and cache versions cap at /3 --
the expected one-time doubling from activation-checkpoint recompute, whose guard failure is
`requires_grad mismatch`). `TORCH_LOGS=graph_breaks` then showed 32 breaks landing exactly on the
`torch._dynamo.disable`d FLA entry points. GLM-5.3-Flash is KDA-dominated (34 of 45 layers), so
every KDA layer broke the surrounding compiled region and what remained compiled was fragments
whose guard and re-entry cost exceeded any fusion benefit.

The workaround was chosen knowingly -- the old comment said as much: "XTuner's own GatedDeltaNet
sidesteps this by owning the op behind a `torch.library.custom_op`; KDA uses FLA's kernel
directly, so mark the call itself as untraceable instead." This takes the GatedDeltaNet route
properly. `xtuner/v1/ops/kda` calls FLA's `chunk_kda_fwd`/`chunk_kda_bwd` behind custom ops with
fake implementations, so dynamo traces through them as opaque nodes and `prepare_chunk_indices`
-- whose `.tolist()` on `cu_seqlens` inductor cannot lower, and which is why FLA's own
`chunk_kda` carries `@torch.compiler.disable` -- stays hidden without a break.

Only the call shape this model uses is supported (no initial_state/output_final_state, gate
precomputed by `fused_kda_gate`, no FLA CP); anything else raises rather than silently taking a
different path. The recurrent kernel keeps its `dynamo.disable`: it only runs at `seq_len <= 64`,
which compiled training never reaches, and the branch is selected from a Python int so only the
taken branch is traced. The short convolution's break remains and is left for a follow-up.

Result (8-GPU, 20 steps, PACK_MAX_LENGTH=16384, identical data): steady-state 1.21s/step and
tgs@20 13381.1, against 5168.4 before this change and 12203.1 for eager -- compile is now ~10%
faster than eager instead of 2.4x slower. Peak memory 82.11 GB. step20 loss 10.31426430 vs
eager's 10.31408882 (1.8e-4), grad_norm 48.39 vs 48.41.

Test Plan:
  - Bitwise parity against `fla.ops.kda.chunk_kda` on a packed two-document batch: forward and
    all five gradients (q/k/v/g/beta) at max|diff| = 0.0. This caught a real bug in the first
    attempt -- FLA feeds the *L2-normed* q/k to `l2norm_bwd`, and using the originals corrupted
    only the q/k gradients (~1.0-1.5 max diff) while the forward stayed bitwise correct.
  - tests/model/test_glm53_kda.py 5/5, including HF `Glm5NextTextLinearAttention` parity.
  - Graph breaks on the compiled text model drop from 32 to 30, with the `chunk_kda` break gone.

Co-Authored-By: Claude Opus 5 <noreply@anthropic.com>
Documents how the `MODEL_COMPILE=1` anomaly in the text matrix was diagnosed -- ruling out config
drift, data drift, warmup amortization (the explanation originally written in the PR, which the
per-step series disproves) and runaway recompilation, before `TORCH_LOGS=graph_breaks` located 32
breaks on the `dynamo.disable`d FLA entry points -- and records the measured before/after.

Co-Authored-By: Claude Opus 5 <noreply@anthropic.com>
Follows the `chunk_kda` port with the two remaining KDA entry points, for the same reason: FLA's
`fused_kda_gate` carries `@torch.compiler.disable` ("Skip calling `torch.compiler.disable()`d
function"), and its causal-conv dispatcher reaches Triton launch helpers dynamo cannot trace, so
each KDA layer broke the enclosing compiled region three more times. Both now call FLA's
`kda_gate_fwd`/`kda_gate_bwd` and `causal_conv1d_fwd`/`causal_conv1d_bwd` behind
`torch.library.custom_op`, so the kernels -- and therefore the numerics -- are unchanged.

`build_pools`' whole-sequence guard reads `cu_seq_lens_q[-1]`, which is a host sync and a break
in every DSA layer. It catches a caller that forgot to gather across the SP mesh -- a programming
error, not a data condition -- so it now runs in eager only, where every test and the first
training step exercise it.

Graph breaks on the compiled text model drop from 30 to 10. **Throughput does not move**
(tgs@20 13376.2 vs 13381.1, steady-state 1.21s/step either way): the whole regression came from
the single `chunk_kda` break, which sat inside the largest compiled region, while these were
cheap. Break count is the wrong metric; where the break falls is what matters. They are still
worth having -- they remove per-layer host syncs and are what would let
`KimiDeltaAttention.forward` run `fullgraph=True` like Qwen3.5's `GatedDeltaNet.forward`.

Only KDA's call shapes are supported (no residual/initial_state/final state for the conv, fp32
gate output); anything else raises rather than silently taking a different path.

Test Plan:
  - Bitwise parity against `fla.ops.kda.gate.fused_kda_gate` and
    `fla.modules.conv.causal_conv1d`: forward and every gradient at max|diff| = 0. Writing that
    check surfaced two real details -- KDA's `dt_bias` is `[H*K]`, not `[H]` (FLA reduces
    `dg.view(-1, H*K).sum(0)`), and `kda_gate_bwd` returns `dg` as `type_as(g)`, which my first
    fake implementation wrongly declared fp32; under compile that would have produced silently
    wrong dtypes.
  - tests/model/test_glm53_kda.py 5/5, test_glm53_decoder_layer.py 7/7, test_glm53_dsa.py 16/16.
  - 20-step 8-GPU run: step20 loss 10.31386757 vs eager's 10.31408882 (2.2e-4).

Co-Authored-By: Claude Opus 5 <noreply@anthropic.com>
Adds the measurement that matters most from this round: removing 20 further graph breaks changed
throughput by nothing, so the whole compile regression came from the single `chunk_kda` break
inside the largest compiled region. Break count is the wrong metric. Records why the last KPool
break is deliberately left alone, and the mechanism-by-mechanism comparison against Qwen3.5-VL.

Co-Authored-By: Claude Opus 5 <noreply@anthropic.com>
@jayhenry
jayhenry force-pushed the feat/glm53flash-f6-text-moe branch from c65bb23 to 1177d37 Compare September 24, 2026 20:22
The compose model asserted a `SequenceContext` never carried both `pixel_values` and
`pixel_values_videos`, citing the tokenize fn's rule that a sample is image-only or video-only.
But at the model a `SequenceContext` is a *pack*, and a pack routinely holds an image sample next
to a video sample. The per-sample rule is already enforced where it belongs, in the tokenize fn;
enforcing it again per pack rejected ordinary training batches.

It never fired before because no video sample had ever trained: until the F1 fix they all became
33-token fakes. Once they trained for real, the first pack to draw both an image sample and a
video sample crashed the run on step 2 with "a mixed-media SequenceContext should never reach the
compose model". Each modality is spliced onto its own positions (`mm_token_type_ids` 1 vs 2), so
the two are independent; both now run when both are present, each with its own global
placeholder-count check.

The VL launcher's `SAMPLE_MAX_LENGTH` default moves from 4096 to 16384 (= `PACK_MAX_LENGTH`).
A visual sample cannot be truncated -- cutting its span corrupts it, so the tokenize fn now drops
it -- and ci_vl's two-video samples expand to ~14k tokens. At 4096 every video sample was filtered
out and the run silently became image-only; the 4096 had been copied from the text launcher,
where it suits alpaca.

Test Plan:
  - `test_pack_with_an_image_sample_and_a_video_sample` replaces the test that asserted the
    rejection: a pack with one image token and one video token, whose logits must equal splicing
    each modality's `get_visual_features` onto its own positions by hand and running the language
    model (rtol=0, atol=0). It fails on the previous code with the exact production error and
    passes with the fix. tests/model/test_glm53_compose.py 6/6.
  - First VL step with real video: 14455 of 16384 tokens used (88%, previously ~50%), with 57k
    image patches.

Co-Authored-By: Claude Opus 5.5 <noreply@anthropic.com>
longboat2010 added a commit to Ascend-SHL-SACT/xtuner that referenced this pull request Sep 26, 2026
A pack that holds both an image sample and a video sample used to enter the
tower once per modality. The tower's FSDP mesh is the world group and the
language tower's is the dp group, so ranks disagreed on the collective
sequence and baseline hung on step 4. One concatenated call keeps every rank
in lockstep, and each grid row stays its own attention segment.

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