Skip to content

[Feature] Add GLM-5.3-Flash F5: NoPE DSA + KPool indexer + clamped SwiGLU - #2108

Open
jayhenry wants to merge 5 commits into
feat/glm53flash-f4-mhcfrom
feat/glm53flash-f5-nope-dsa
Open

jayhenry wants to merge 5 commits into
feat/glm53flash-f4-mhcfrom
feat/glm53flash-f5-nope-dsa

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 ← you are here
  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

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


Summary

Stack layer 4/7 of GLM-5.3-Flash support (base: layer 3, F4 mHC).

Implements the F5 milestone: the NoPE (qk_rope_head_dim=0) DeepSeek Sparse Attention layers, their KPool indexer (pools of index_kpool consecutive tokens scored together instead of per-token top-k), a new flash_mla_cudnn SparseMLA backend (FlashMLA forward + existing cuDNN backward), and the clamped SwiGLU activation used by GLM-5.3-Flash's dense/shared/MoE MLPs.

Key pieces:

  • xtuner/v1/ops/act_fn.py: native_clamped_swiglu + MoEActFnConfig support, wired through DenseMLP/MoEMLP.
  • xtuner/v1/ops/sparse_mla/kpool.py: pool layout, causal ranges, and top-k pool selection (torch reference + TileLang-backed production path).
  • xtuner/v1/ops/sparse_mla/flash_mla_cudnn.py: new SparseMLA backend.
  • xtuner/v1/model/moe/glm53/nope_dsa_mla.py: NoPEDSAMLAConfig, KPoolIndexer, NoPEDSAMultiLatentAttention.
  • xtuner/v1/ops/sparse_mla/protocol.py: KPoolIndexerBackend/KPoolTopKIndicesProtocol (GLM-5.3-Flash's KPool only supports torch/tilelang, unlike GLM-5.2's 6-way DSAIndexerBackend).

Also folds in what were originally two separate follow-up commits, now part of this PR:

  • DSA backends get explicit defaults instead of inheriting: indexer_backend used to default to None and fall back to sparse_mla_backend, but the two name different, only partially-overlapping backend vocabularies (flash_mla_cudnn is SparseMLA-only; deep_gemm_fp8/cute_dsl are indexer-only), so the fallback could hand the indexer factory a backend it has no implementation for. Both DSAMLAConfig/NoPEDSAMLAConfig fields now default to tilelang directly and explicitly.
  • Bound the KPool indexer's logits tile with query chunking, mirroring the existing tilelang DSA selector's query-chunk bound.

Known gaps recorded in doc/progress.md: tilelang sparse_mla_backend and deep_gemm_fp8 indexer_backend are not implemented for NoPE (explicit NotImplementedError, not silent fallback); KPool/NoPE-DSA sequence-parallel path is written per design but not yet GPU-tested at sp_size>1.

Test Plan

tests/model/test_glm53_dsa.py (12), test_glm53_nope_dsa_mla.py (2), test_flash_mla_cudnn_sparse_mla.py (3) all pass; GLM-5.2's existing tests/module/attention/test_dsa_mla.py (15) rerun clean against the shared tilelang.py/dsa_mla.py edits, confirming no regression. Numerical oracle: transformers 5.17.0's glm5_next model.

self.up_proj = build_linear(self.hidden_size, self.intermediate_size, bias=mlp_bias, float8_cfg=float8_cfg)
self.down_proj = build_linear(self.intermediate_size, self.hidden_size, bias=mlp_bias, float8_cfg=float8_cfg)
self.act_fn = get_act_fn(hidden_act)
self.act_fn = get_gated_act_fn(hidden_act, swiglu_limit)

Copy link
Copy Markdown
Collaborator Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

去掉 get_gated_act_fn,使用 MoEActFnConfig ,是不是更好的方案?

from .flash_mla import flash_mla_sparse_mla

return flash_mla_sparse_mla
if backend == "flash_mla_cudnn":

Copy link
Copy Markdown
Collaborator Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

cudnn_dsa 已经改为前向 flash mla,去掉这个分支?

starts,
ends,
select_k,
query_chunk_size=query_chunk_size,

@jayhenry jayhenry Oct 9, 2026 •

Copy link
Copy Markdown
Collaborator Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

增加 selector 参数,以便支持 tilelang_deepselect

def get_kpool_topk_indices(backend: KPoolIndexerBackend) -> KPoolTopKIndicesProtocol:
if backend == "torch":
return torch_kpool_topk_indices
if backend == "tilelang":

Copy link
Copy Markdown
Collaborator Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

支持 tilelang_deepselect

jayhenry and others added 5 commits October 10, 2026 18:37
…iGLU

Implements the F5 milestone from doc/xtuner_glm5p3flash_design.md: the
NoPE (qk_rope_head_dim=0) DeepSeek Sparse Attention layers, their KPool
indexer (pools of index_kpool consecutive tokens scored together instead
of per-token top-k), a new flash_mla_cudnn SparseMLA backend (FlashMLA
forward + existing cuDNN backward), and the clamped SwiGLU activation
used by GLM-5.3-Flash's dense/shared/MoE MLPs.

- xtuner/v1/ops/act_fn.py: native_clamped_swiglu + MoEActFnConfig support.
- xtuner/v1/module/decoder_layer/{dense,moe}_decoder_layer.py: wire
  swiglu_limit through DenseMLP/MoEMLP.
- xtuner/v1/ops/sparse_mla/kpool.py: pool layout, causal ranges, and
  top-k pool selection (torch reference + TileLang-backed production
  path, reusing the existing indexer kernel unmodified since relu's
  homogeneity makes head_dim^-0.5 movable from the relu argument into
  the per-head weight without changing the result).
- xtuner/v1/ops/sparse_mla/flash_mla_cudnn.py: new SparseMLA backend.
- xtuner/v1/ops/sparse_mla/tilelang.py: widen the hardcoded 576 head-dim
  check to a (head_dim, value_dim) whitelist.
- xtuner/v1/model/moe/glm52/dsa_mla.py: reject flash_mla_cudnn as an
  indexer_backend (SparseMLABackend widened, but this backend has no
  indexer counterpart).
- xtuner/v1/model/moe/glm53/nope_dsa_mla.py: NoPEDSAMLAConfig,
  KPoolIndexer, NoPEDSAMultiLatentAttention.
- xtuner/v1/ops/sparse_mla/protocol.py: KPoolIndexerBackend (restricts
  GLM-5.3-Flash's KPool indexer_backend to "torch"/"tilelang" -- unlike
  GLM-5.2's 6-way DSAIndexerBackend, there's no cudnn_dsa/flash_mla/
  deep_gemm_fp8/cute_dsl KPool kernel) and KPoolTopKIndicesProtocol.
- xtuner/v1/ops/sparse_mla/__init__.py: get_kpool_topk_indices(backend),
  mirroring get_dsa_topk_indices's style -- explicitly raises for any
  backend other than "torch"/"tilelang" instead of silently falling
  through. KPoolIndexer resolves this once in __init__ into
  self._topk_indices_fn instead of re-branching on every forward call.
- xtuner/v1/model/moe/glm53/nope_dsa_mla.py: NoPEDSAMultiLatentAttention.
  forward branches on freeze_dsa_indexer (torch.no_grad() only when
  frozen), mirroring GLM-5.2's per-token indexer, ahead of a future
  differentiable indexer output -- today's kpool_topk_indices/
  torch_kpool_topk_indices still only ever return an int32 index
  tensor, so this doesn't yet change what's trainable (confirmed via a
  freeze_dsa_indexer=False smoke run: indexer params get
  requires_grad=True but no actual gradient). Calls self.indexer(...)
  directly instead of through reuse_during_recompute, which retained
  topk_ids' activations across the backward recompute pass to avoid
  recomputing them; dropped since nothing currently offloads or reuses
  that retained memory.

Known gaps recorded in doc/progress.md: tilelang sparse_mla_backend and
deep_gemm_fp8 indexer_backend are not implemented for NoPE (explicit
NotImplementedError, not silent fallback); KPool/NoPE-DSA sequence-
parallel path is written per design but not yet GPU-tested at sp_size>1;
whether index topk ids can be offloaded is still an open question, not
attempted here.

Test plan: tests/model/test_glm53_dsa.py (12), test_glm53_nope_dsa_mla.py
(2), test_flash_mla_cudnn_sparse_mla.py (3) all pass; GLM-5.2's existing
tests/module/attention/test_dsa_mla.py (15) rerun clean against the
shared tilelang.py/dsa_mla.py edits, confirming no regression.
Numerical oracle: transformers 5.17.0's glm5_next model. Two near-tied
top-k floating-point sensitivities were root-caused via seed sweeps and
a targeted trace script (not worked around by loosening tolerances
blindly) -- see doc/progress.md F5 section for the full analysis.

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

Also folds in what were originally two follow-up commits, now part of F5 itself
rather than separate history:
- Give the DSA backends explicit defaults instead of inheriting: both
  DSAMLAConfig/NoPEDSAMLAConfig fields now default to tilelang directly,
  with indexer_backend never falling back to sparse_mla_backend (the two
  name different, only partially overlapping backend vocabularies).
- Bound the KPool indexer's logits tile with query chunking, mirroring the
  existing tilelang DSA selector's query-chunk bound.
* [Fix] Define and forward the TileLang DSA indexer selector

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.

* [Fix] Guard NoPE DSA config against unsupported settings at config time

Two config-time guards for NoPEDSAMLAConfig, closing review findings #1
and #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.

---------

Co-authored-by: liutong <liutong@pjlab.org.cn>
b03bf01/fe0037f0 switched _cudnn_dsa_sparse_mla_backward_op from the
log2-space LSE contract to natural-log LSE, but this backend kept
forwarding the log2 value, so every GLM-5.3 training silently fed the
cuDNN backward a 1.4427x-inflated LSE. Forward is unaffected; backward
gradients come out systematically shrunk (grad_norm ~20-30% low) and
loss drifts (+0.043 over 300 steps vs the pre-change baseline).
Save and pass the natural-log softmax_lse instead.

Co-authored-by: liutong <liutong@pjlab.org.cn>
The tiny GLM-5.2 cases selected TileLang dimensions and sparse MLA kernels outside their supported shape set. Keep the actual indexer path while using a supported index dimension and the torch sparse MLA reference. The F5 routed activation check now uses the F5 public config rather than importing the F6 text model.
The colocate test inherited WORLD_SIZE from other cases and requested 32 workers on an 8-GPU node. Direct sampled KL is noisy for greedy rollouts, so retain its finite check and bound the stable K3 metric. Check temporary CUDA tensor lifetime directly instead of requiring allocator reserved bytes to decrease.
@jayhenry
jayhenry force-pushed the feat/glm53flash-f5-nope-dsa branch from eb6ad8f to 5344d1d Compare October 10, 2026 19:01

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.

2 participants