[muon] Per-head Muon for linear-attention layers: read head geometry from the owning module (#8436)
Closes #8420. Stacked on #8384 — that branch is the base, so the diff to
review is the single commit 622c28de3; the rest is #8384.
@delock answered the two model questions in #8420: Muon Split applies to
KDA, and not to GLM's indexer. This implements exactly that.
## The problem
`#8384` derives candidate geometries from the config: head counts
through `AutoTPMeta`, per-head widths from `head_dim` or the MLA fields.
Kimi-K3's linear-attention layers are built from `linear_attn_config`
instead:
```python
# modeling_kimi_k3_linear.py, KimiDeltaAttention
self.head_dim = config.linear_attn_config["head_dim"] # 32
self.num_heads = config.linear_attn_config["num_heads"] # 8
self.q_proj = nn.Linear(self.hidden_size, self.head_k_dim * self.num_k_heads, bias=False)
```
`q_proj` is (256, 1024) = 8 x 32. Nothing at the top level says 32:
`head_dim` is 74 and `qk_nope + qk_rope` is 96, so every candidate
failed the width check. The head count was right and the width was
wrong, so the projections were declined rather than recognized.
## The change
Map the parameter to the module that owns it and ask that module for the
geometry it was built with. It is one more candidate, not an override —
`#8384`'s shape confirmation still decides, and a module that disagrees
with the config on the head count is still ambiguous and still skipped.
Only modules that say they are attention are asked. That is what keeps
the indexer out: `GlmMoeDsaIndexer.wq_b` is (512, 512) and
`index_n_heads 8 * index_head_dim 64` is exactly 512, so reading
`n_heads`/`head_dim` off any module carrying them would tag it. It is
declined by the module, not by the shape, and there is a test asserting
the geometry does confirm so that the reason stays honest.
The cost of that choice is stated in a test: Qwen3-Next's
`Qwen3NextGatedDeltaNet` does not say attention, so it keeps the
full-matrix path. That is `#8384`'s behaviour rather than a regression,
and widening it is one marker. The alternative — ask every module, name
the ones to skip — tags the indexer by default, which is the opposite of
what #8420 concluded.
## Both models, instantiated
`inference-optimization/Kimi-K3-0.40B`, 318 Muon parameters:
| owner | parameter | shape | before | after |
| --- | --- | --- | --- | --- |
| `KimiMLAAttention` (2 layers) | `q_b_proj` | (768, 256) | 8 heads
(`mla-q`) | unchanged |
| `KimiMLAAttention` (2 layers) | `kv_b_proj` | (1024, 128) | 8 heads
(`mla-kv`) | unchanged |
| `KimiDeltaAttention` (6 layers) | `q_proj` | (256, 1024) | full matrix
| **8 heads (`owner-q`)** |
| `KimiDeltaAttention` (6 layers) | `k_proj` | (256, 1024) | full matrix
| **8 heads (`owner-k`)** |
| `KimiDeltaAttention` (6 layers) | `v_proj` | (256, 1024) | full matrix
| **8 heads (`owner-v`)** |
| `KimiDeltaAttention` | `g_proj` | (256, 1024) | full matrix |
unchanged |
| `KimiDeltaAttention` | `f_b_proj`, `b_proj`, `o_proj` | — | full
matrix | unchanged |
4 tagged -> 22. `g_proj` is the case worth noting: same 256 x 1024 shape
as `q_proj`, in the same module, and it stays whole because its leaf
name is not a projection name. There is a test for it.
`inference-optimization/GLM-5.2-0.8B-A0.8B`, 69 Muon parameters —
unchanged at 12:
```
TAGGED: q_b_proj (4096, 512) x6 -> 16 (mla-q)
kv_b_proj (5120, 128) x6 -> 16 (mla-kv)
DECLINED: wq_b (512, 512) x3 width-mismatch
wk (64, 2048) x3 width-mismatch
```
## End to end, real checkpoints
bf16, ZeRO-1, `per_head_muon: true`, 6 steps, loaded weights rather than
`from_config` (the KDA layer has an uninitialized `dt_bias`, so a
randomly initialized hybrid gives NaN on step 0 regardless of the
optimizer):
```
Kimi-K3-0.40B tagged=22 (18 KDA + 4 MLA) loss 20.00 -> 18.25 all params finite
GLM-5.2-0.8B tagged=12 indexer_tagged=0 loss 12.29 -> 9.98 all params finite
```
## Tests
12 cases added to
`tests/unit/runtime/zero/test_per_head_muon_tagging.py`, 54 passing.
Reverting the code change fails exactly the four that assert the new
behaviour:
```
FAILED test_linear_attention_geometry_comes_from_the_owning_module[q_proj] - assert None == 8
FAILED test_linear_attention_geometry_comes_from_the_owning_module[k_proj] - assert None == 8
FAILED test_linear_attention_geometry_comes_from_the_owning_module[v_proj] - assert None == 8
FAILED test_the_value_projection_uses_the_value_width - assert None == 8
```
The other eight pass both before and after by design — they pin that
nothing widened: the indexer, a module that does not say attention,
geometry that does not match the shape, the non-projection matrices of a
KDA layer, and a standard Llama where the config and the module agree.
The two existing tests that assert Kimi's KDA is declined on a
config-only model are unchanged and still pass. A `SimpleNamespace`
config has no modules, so they now document the narrower thing they were
always testing: the config alone cannot describe this layout.
Whether per-head helps on linear-attention heads is a separate question
from whether the split is well-defined on them, and #8384 has the
measurements on that.
---------
Signed-off-by: alanhuangyoo <alanhuangyoo@gmail.com>
Co-authored-by: Claude Opus 5 <noreply@anthropic.com>
Co-authored-by: pengdurice <pengduhit@gmail.com>
Co-authored-by: Ma, Guokai <guokai.ma@gmail.com>