kempnerforge.model.moe

Mixture-of-Experts feed-forward layer for KempnerForge models.

Functions

build_moe(dim, hidden_dim, num_experts, top_k)

Build an MoE layer, composing router + experts from Registry.

grouped_expert_forward(x_sorted, ...)

Run every expert on its own token group with ragged grouped GEMMs.

grouped_expert_forward_packed(x_sorted, ...)

Same as grouped_expert_forward over pre-packed (E, in, out) weights.

scale_by_expert_load(expert_out, ...)

Rescale each expert's rows (sorted by expert) by avg_tokens / its token count.

Classes

MoEMLP

Mixture-of-Experts feed-forward layer.

kempnerforge.model.moe.grouped_expert_forward(x_sorted, tokens_per_expert, experts)[source]

Run every expert on its own token group with ragged grouped GEMMs.

Each expert sees exactly its tokens – nothing is padded to the busiest expert – so memory and time follow the total number of routed tokens, not the imbalance.

Parameters:
  • x_sorted (torch.Tensor) – (total_tokens, dim) token features sorted by expert index.

  • tokens_per_expert (torch.Tensor | Sequence[int]) – (E,) token count per expert, in expert order.

  • experts (nn.ModuleList) – Expert modules whose weights are stacked for the grouped GEMM.

Returns:

(total_tokens, dim) expert outputs in the same sorted order as the input.

Return type:

torch.Tensor

kempnerforge.model.moe.grouped_expert_forward_packed(x_sorted, tokens_per_expert, up_w, down_w, gate_w, activation)[source]

Same as grouped_expert_forward over pre-packed (E, in, out) weights.

Parameters:
  • x_sorted (torch.Tensor) – (total_tokens, dim) token features sorted by expert index.

  • tokens_per_expert (torch.Tensor | Sequence[int]) – (E,) token count per expert, in expert order.

  • up_w (torch.Tensor) – (E, dim, hidden) packed up-projection weights.

  • down_w (torch.Tensor) – (E, hidden, dim) packed down-projection weights.

  • gate_w (torch.Tensor | None) – (E, dim, hidden) packed gate weights for SwiGLU, else None.

  • activation – Applied to the up-projection when gate_w is None.

Returns:

(total_tokens, dim) expert outputs in the same sorted order as the input.

Return type:

torch.Tensor

kempnerforge.model.moe.scale_by_expert_load(expert_out, tokens_per_expert, num_experts)[source]

Rescale each expert’s rows (sorted by expert) by avg_tokens / its token count.

Parameters:
Return type:

torch.Tensor

class kempnerforge.model.moe.MoEMLP[source]

Bases: Module

Mixture-of-Experts feed-forward layer.

Composes a router (from “router” registry) with N expert MLPs (from “mlp” registry). Drop-in replacement for dense MLP — same forward signature.

Stores aux_loss after each forward for collection by the training loop.

__init__(router, experts, shared_expert=None, capacity_factor=0.0, gradient_scale=False, packed_experts=False)[source]
Parameters:
  • router (nn.Module)

  • experts (nn.ModuleList)

  • shared_expert (nn.Module | None)

  • capacity_factor (float)

  • gradient_scale (bool)

  • packed_experts (bool)

Return type:

None

property aux_loss: torch.Tensor
property z_loss: torch.Tensor
property expert_counts: torch.Tensor
forward(x)[source]

Forward pass dispatching tokens to experts.

Parameters:

x (torch.Tensor) – (batch, seq_len, dim)

Returns:

(batch, seq_len, dim)

Return type:

torch.Tensor

kempnerforge.model.moe.build_moe(dim, hidden_dim, num_experts, top_k, activation='silu', router_type='softmax_topk', shared_experts=0, capacity_factor=0.0, gradient_scale=False, sequence_aux_loss_weight=0.0, bias_schedule='constant', packed_experts=False)[source]

Build an MoE layer, composing router + experts from Registry.

Parameters:
  • dim (int) – Model dimension.

  • hidden_dim (int) – Expert FFN hidden dimension.

  • num_experts (int) – Number of routed experts.

  • top_k (int) – Experts selected per token.

  • activation (str) – MLP activation (registry key).

  • router_type (str) – Router registry key.

  • shared_experts (int) – Number of shared experts (always active).

  • capacity_factor (float) – Token capacity per expert (0=unlimited, >0=cap).

  • gradient_scale (bool) – Per-expert gradient normalization.

  • sequence_aux_loss_weight (float) – Sequence-level balance loss weight (sigmoid router only).

  • bias_schedule (str) – Bias update rate schedule (sigmoid router only).

  • packed_experts (bool) – Pack expert weights into one tensor per projection.

Return type:

MoEMLP