Skip to content

[Fix] Define and forward the TileLang DSA indexer selector - #2143

Merged
jayhenry merged 2 commits into
InternLM:feat/glm53flash-f5-nope-dsafrom
Denny991:liutong/fix/tilelang-dsa-indexer-selector
Oct 9, 2026
Merged

jayhenry merged 2 commits into
InternLM:feat/glm53flash-f5-nope-dsafrom
Denny991:liutong/fix/tilelang-dsa-indexer-selector

Conversation

@Denny991

@Denny991 Denny991 commented Oct 9, 2026 •

Copy link
Copy Markdown
Collaborator

What

Three small fixes for the GLM-5.3-Flash f5 stack, found during the stack review.

1. (Blocking) Define and forward the TileLang DSA indexer selector

3de92e7b ("[Feature] Add DeepSelect top-k selector for the TileLang DSA indexer") made tilelang_indexer_topk_from_ranges dispatch on a selector variable without declaring it in the signature, so every call raises NameError: name 'selector' is not defined at sparse_mla/tilelang.py:250.

GLM-5.3-Flash's default KPool indexer backend (indexer_backend="tilelang", nope_dsa_mla.py) routes through this wrapper, so on the current stacked branch (pr-2111 head a537fcb4) no GLM-5.3 text training run can start a single step. Reproduced on 8xH200 with the aligned GLM-5.3-Flash reduced-layer 16k SFT config; all 8 ranks crash at step 1:

[rank6]:     _tilelang_dsa_topk_indices_from_ranges if selector == "torch" else ...
[rank6]:                                               ^^^^^^^^
[rank6]: NameError: name 'selector' is not defined

Fix: declare selector: Literal["torch", "deep_select"] = "torch" in the wrapper (default preserves the pre-DeepSelect behavior for existing callers such as the KPool indexer), and forward it from tilelang_dsa_topk_indices -- without the forwarding, the tilelang_deepselect backend silently ran the torch.topk kernel instead of DeepSelect's radix-select kernel.

2. Reject freeze_dsa_indexer=False in NoPEDSAMLAConfig.build()

Mirrors GLM-5.2's guard. The KPool indexer only returns int32 indices and its gradients are cut by no_grad, so the flag never trains anything; it silently leaves indexer params requires_grad=True, paying optimizer state and activation memory for parameters that cannot update, with no error.

3. Require num_attention_heads % 64 == 0 for flash_mla_cudnn at config time

The production default backend previously checked FlashMLA head alignment only at the first forward (flash_mla_cudnn.py); the model validator now rejects misaligned configs immediately. The published GLM-5.3-Flash config (64 heads) is unaffected; the torch reference backend keeps no alignment requirement. The constant comes from flash_mla.py's _FLASH_MLA_HEAD_ALIGNMENT for a single source of truth.

(Review finding #2, KPool gate init zeros-vs-ones, is deliberately not addressed here: it only affects from-scratch training and the current init is a documented, defensible choice -- proposing offline discussion instead.)

Tests

  • tests/ops/test_tilelang_dsa_indexer_selector.py (new): CPU routing tests for both selector values and the default, plus a CUDA-gated test that tilelang_dsa_topk_indices(selector="deep_select") reaches the DeepSelect kernel. Either selector regression would have been caught by these.
  • tests/model/test_glm53_nope_dsa_mla.py: new TestNoPEDSAMLAConfigGuards -- both config rejections plus the torch-backend exemption.

Verification (8xH200)

  • Unit tests: test_glm53_dsa.py + test_glm53_decoder_layer.py + test_glm53_nope_dsa_mla.py + test_tilelang_dsa_indexer_selector.py -- 33 passed, 2 failed; both failures fail identically on the pristine f5 base and are unrelated (one is an inductor env limitation).
  • E2E: GLM-5.3-Flash reduced-layer 16k SFT, aligned config. Before: crashes at step 1 with the NameError above. After: 20/20 steps at ~0.93 s/step, and step-1 reduced_llm_loss = 10.71162510 -- bit-identical to the healthy pre-regression baseline.

Found during GLM-5.3-Flash stack review; targeting the f5 stack branch since all three findings trace to it.

liutong added 2 commits October 9, 2026 15:53
3de92e7 introduced the DeepSelect selector switch but referenced
`selector` inside tilelang_indexer_topk_from_ranges without declaring
it in the signature, so every call raised NameError. GLM-5.3's default
KPool indexer backend (indexer_backend="tilelang") routes through this
wrapper, so training could not start a single step on the f5 stack.

The outer tilelang_dsa_topk_indices also accepted a selector but dropped
it when calling the wrapper, so the tilelang_deepselect backend silently
ran the torch.topk kernel instead of DeepSelect's radix-select kernel.

- Declare selector in tilelang_indexer_topk_from_ranges, defaulting to
  "torch" to preserve the pre-DeepSelect behavior for existing callers
  (the KPool indexer).
- Forward selector from tilelang_dsa_topk_indices.
- Add CPU routing tests for the wrapper plus a CUDA-gated forwarding
  test that would have caught both regressions.
Two config-time guards for NoPEDSAMLAConfig, closing review findings InternLM#1
and InternLM#3 from the GLM-5.3-Flash stack review:

- freeze_dsa_indexer=False now raises ValueError in build(), mirroring
  GLM-5.2's guard. The indexer only returns int32 indices and its
  gradients are cut by no_grad, so the flag silently left indexer
  params requires_grad=True -- paying optimizer state and activation
  memory for parameters that never train, with no error.
- sparse_mla_backend='flash_mla_cudnn' (the production default) now
  requires num_attention_heads % 64 == 0 in the model validator,
  instead of failing at the first forward inside the FlashMLA kernel.
  The production config (64 heads) is unaffected; the torch reference
  backend keeps no alignment requirement.

Also adds three tests covering both rejections and the torch-backend
exemption.
@jayhenry
jayhenry merged commit 9c1f93c into InternLM:feat/glm53flash-f5-nope-dsa Oct 9, 2026
1 of 3 checks passed
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.

2 participants