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:
@@ -88,8 +88,26 @@ def test_preset_0_5b_schedule():
|
||||
assert cfg.layer_specs()[3] == ("mla", "moe")
|
||||
|
||||
|
||||
def _dense_moe_forward(moe: LatentMoE, x: torch.Tensor) -> torch.Tensor:
|
||||
"""Old dense path: run every routed expert, then gather top-k."""
|
||||
z = moe.down(x)
|
||||
logits = moe.router(x)
|
||||
topk = torch.topk(logits, moe.top_k, dim=-1)
|
||||
ids = topk.indices
|
||||
probs = F.softmax(topk.values, dim=-1)
|
||||
all_out = torch.stack([expert(z) for expert in moe.experts])
|
||||
B, T, _ = x.shape
|
||||
all_out = all_out.permute(1, 2, 0, 3).reshape(B * T, moe.n_routed, moe.latent_size)
|
||||
u = z.new_zeros(B, T, moe.latent_size)
|
||||
for i in range(moe.top_k):
|
||||
idx = ids[:, :, i].reshape(B * T)
|
||||
sel = all_out[torch.arange(B * T, device=x.device), idx]
|
||||
u = u + probs[:, :, i : i + 1] * sel.reshape(B, T, moe.latent_size)
|
||||
shared = torch.stack([expert(x) for expert in moe.shared]).sum(0)
|
||||
return shared + moe.up(moe.norm(u))
|
||||
|
||||
|
||||
def test_moe_router_activates_topk_only():
|
||||
from kda.layers.latent_moe import LatentMoE
|
||||
torch.manual_seed(3)
|
||||
moe = LatentMoE(hidden_size=32, latent_size=16, n_routed=8, top_k=2, n_shared=1, d_ff=24)
|
||||
x = torch.randn(2, 6, 32)
|
||||
@@ -112,6 +130,62 @@ def test_moe_router_activates_topk_only():
|
||||
torch.testing.assert_close(y, expected_y, atol=1e-5, rtol=1e-5)
|
||||
assert moe.last_route_ids is not None
|
||||
assert moe.last_route_ids.shape[-1] == moe.top_k
|
||||
counts = torch.bincount(moe.last_route_ids.reshape(-1), minlength=moe.n_routed)
|
||||
assert moe.last_capacity == int(counts.max())
|
||||
|
||||
|
||||
def test_moe_sparse_matches_dense_fwd_bwd():
|
||||
torch.manual_seed(3)
|
||||
moe = LatentMoE(hidden_size=32, latent_size=16, n_routed=8, top_k=2, n_shared=1, d_ff=24)
|
||||
x = torch.randn(2, 6, 32)
|
||||
y_sparse = moe(x)
|
||||
y_dense = _dense_moe_forward(moe, x)
|
||||
torch.testing.assert_close(y_sparse, y_dense, atol=1e-5, rtol=1e-5)
|
||||
|
||||
moe.zero_grad(set_to_none=True)
|
||||
xs = x.detach().requires_grad_(True)
|
||||
moe(xs).square().mean().backward()
|
||||
grads_s = {
|
||||
name: param.grad.detach().clone()
|
||||
for name, param in moe.named_parameters()
|
||||
if param.grad is not None
|
||||
}
|
||||
dx_s = xs.grad.detach().clone()
|
||||
|
||||
moe.zero_grad(set_to_none=True)
|
||||
xd = x.detach().requires_grad_(True)
|
||||
_dense_moe_forward(moe, xd).square().mean().backward()
|
||||
grads_d = {
|
||||
name: param.grad.detach().clone()
|
||||
for name, param in moe.named_parameters()
|
||||
if param.grad is not None
|
||||
}
|
||||
dx_d = xd.grad.detach().clone()
|
||||
|
||||
torch.testing.assert_close(dx_s, dx_d, atol=1e-5, rtol=1e-5)
|
||||
assert grads_s.keys() == grads_d.keys()
|
||||
for name in grads_s:
|
||||
torch.testing.assert_close(grads_s[name], grads_d[name], atol=1e-5, rtol=1e-5)
|
||||
|
||||
|
||||
def test_moe_unselected_experts_have_zero_grad():
|
||||
torch.manual_seed(0)
|
||||
moe = LatentMoE(hidden_size=32, latent_size=16, n_routed=8, top_k=2, n_shared=1, d_ff=24)
|
||||
with torch.no_grad():
|
||||
moe.router.weight.zero_()
|
||||
moe.router.weight[0] = 1.0
|
||||
moe.router.weight[1] = 0.5
|
||||
x = torch.ones(2, 4, 32)
|
||||
moe(x).square().mean().backward()
|
||||
selected = set(moe.last_route_ids.reshape(-1).tolist())
|
||||
assert selected == {0, 1}
|
||||
for idx, expert in enumerate(moe.experts):
|
||||
for param in (expert.w_g.weight, expert.w_u.weight, expert.w_o.weight):
|
||||
assert param.grad is not None
|
||||
if idx in selected:
|
||||
assert param.grad.abs().sum() > 0
|
||||
else:
|
||||
assert torch.equal(param.grad, torch.zeros_like(param.grad))
|
||||
|
||||
|
||||
def test_k3_causal_future_does_not_change_past_logits():
|
||||
|
||||
Reference in New Issue
Block a user