class MoELayer(nnx.Module):
"""Mixture-of-experts FFN: a token-choice top-k router over E experts."""
def __init__(self, num_hiddens, num_experts, num_active, rngs=None):
rngs = nnx.Rngs(0) if rngs is None else rngs
self.num_experts, self.num_active = num_experts, num_active
self.router = nnx.Linear(num_hiddens, num_experts, use_bias=False,
rngs=rngs)
self.experts = nnx.List([d2l.FeedForward(num_hiddens, rngs=rngs)
for _ in range(num_experts)])
self.expert_bias = nnx.Variable(jnp.zeros(num_experts))
self.usage = nnx.Variable(jnp.zeros(num_experts))
self.aux_loss = nnx.Variable(jnp.zeros(()))
def __call__(self, X):
probs = jax.nn.softmax(self.router(X), -1) # (B, T, E)
scores = probs + self.expert_bias[...] # selection only
_, idx = jax.lax.top_k(scores, self.num_active) # (B, T, k)
mask = jax.nn.one_hot(idx, self.num_experts).sum(-2)
gates = probs * mask # weight = p_i
Y = jnp.stack([e(X) for e in self.experts], -1) # (B, T, d, E)
out = (Y * gates[..., None, :]).sum(-1)
frac = mask.sum((0, 1)) / mask.sum() # realized load
self.usage[...] = self.usage[...] + mask.sum((0, 1))
self.aux_loss[...] = self.num_experts * (
frac * probs.mean((0, 1))).sum()
return out