LatentMoE: sparse permute-dispatch + padded bmm

Replace dense all-expert forward (16 experts × all tokens) with
permute-dispatch: sort token-expert pairs by expert id, pad to
[R, C, ℓ] (C = max tokens per expert), run 3 bmm calls for the
batched SiTU-GLU activation, then scatter-add weighted results back.

Routed expert FLOPs drop from R·N to R·C (C ≈ N·k/R under uniform
routing). SiTU parameter structure unchanged; checkpoint compatible.

Tests: sparse-vs-dense fwd/bwd equivalence, unselected expert zero
grad, last_capacity tracking.
This commit is contained in:
dela
2026-08-25 17:47:44 +08:00
parent 584f7e9e73
commit d1da0816f2
3 changed files with 131 additions and 11 deletions
+1 -1
View File
@@ -73,7 +73,7 @@ class K3Config:
if name == "toy":
return cls()
if name in {"0.5b", "500m"}:
# H * head_dim == hidden. Routed 16: LatentMoE still runs every expert.
# H * head_dim == hidden. Routed 16 Top-2; LatentMoE padded bmm.
# ~482M with tied Qwen3 embeddings. 6×(3 KDA + 1 MLA).
return cls(
hidden_size=768,