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:
@@ -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,
|
||||
|
||||
Reference in New Issue
Block a user