nemo_automodel.components.moe.mok_experts
nemo_automodel.components.moe.mok_experts
Mixture-of-Kittens expert backend for AutoModel’s shared MoE component.
Module Contents
Classes
Functions
Data
API
Bases: Module
AutoModel expert parameters executed by Mixture-of-Kittens.
Routed parameters use MoK-native contiguous layouts during training while state-dict hooks preserve AutoModel’s established combined expert keys.
Return None because MoK does not support expert bias.
Return AutoModel’s virtual down-projection state-dict tensor.
Return AutoModel’s virtual combined GateUp state-dict tensor.
Return None because MoK does not support expert bias.
Load established AutoModel expert tensors into MoK-native parameters.
Save virtual AutoModel expert tensors instead of MoK-native parameters.
Reject MoE variants outside MoK’s current fused-kernel contract.
Execute fused shared and routed SwiGLU experts.
Parameters:
BF16 tensor of shape [tokens, hidden].
Tensor of shape [tokens, activated_experts].
Int64 tensor of shape [tokens, activated_experts].
BF16 tensor of shape [expert_intermediate, hidden].
BF16 tensor of shape [expert_intermediate, hidden].
BF16 tensor of shape [hidden, expert_intermediate].
Returns: torch.Tensor
BF16 tensor of shape [tokens, hidden] containing the fused shared and
Initialize MoK on an expert-parallel device mesh.
Parameters:
One-dimensional device mesh named ep.
Initialize MoK-native weights from the canonical expert layout.
Draw random values in the same tensor shapes and order as
GroupedExperts before transposing them into MoK’s native
layouts. Otherwise switching only the dispatcher changes a
random-initialized model even when every process uses the same seed.
Parameters:
Device on which initialization kernels execute.
Standard deviation of the normal weight initialization.
Bases: Function
Connect MoK’s explicit backward API to PyTorch autograd.
Run the fused MoK backward pass.
Parameters:
Autograd-owned context containing the forward schedule and tensors.
BF16 tensor of shape [tokens, hidden].
Returns: torch.Tensor | None
Gradients matching every forward input. The runtime and integer
Run a fused MoK forward pass.
Parameters:
Autograd-owned context.
Initialized MoK runtime for this layer’s EP group.
Contiguous BF16 tensor of shape [tokens, hidden].
Contiguous FP32 tensor of shape [tokens, activated_experts].
Contiguous int64 tensor of shape [tokens, activated_experts].
Contiguous BF16 tensor of shape [expert_intermediate, hidden].
Contiguous BF16 tensor of shape [expert_intermediate, hidden].
Contiguous BF16 tensor of shape [hidden, expert_intermediate].
Contiguous BF16 tensor of shape [local_experts, expert_intermediate, hidden].
Contiguous BF16 tensor of shape [local_experts, expert_intermediate, hidden].
Contiguous BF16 tensor of shape [local_experts, hidden, expert_intermediate].
Returns: torch.Tensor
BF16 tensor of shape [tokens, hidden].
Own one layer’s MoK configuration and expert-parallel process group.
Prequantize all routed-weight layouts once per optimizer generation.
Parameters:
Contiguous BF16 tensor of shape [local_experts, expert_intermediate, hidden].
Contiguous BF16 tensor of shape [local_experts, expert_intermediate, hidden].
Contiguous BF16 tensor of shape [local_experts, hidden, expert_intermediate].
Returns: tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor]
Gate, up, and down projection layouts. Each entry contains the normal
Return MoK’s workspace for one local activation shape.
Parameters:
Contiguous BF16 tensor of shape [tokens, hidden].
Number of activated experts per token.
Returns: object
Opaque MoK workspace shared by schedule construction and kernels.
Advance the BF16-weight generation and drop its quantized tensors.
Run MoK’s manual backward.
Parameters:
Opaque schedule returned by forward.
Opaque activation context returned by forward.
Contiguous BF16 tensor of shape [tokens, hidden].
Contiguous BF16 tensor of shape [tokens, hidden].
Contiguous FP32 tensor of shape [tokens, activated_experts].
Contiguous BF16 tensor of shape [expert_intermediate, hidden].
Contiguous BF16 tensor of shape [expert_intermediate, hidden].
Contiguous BF16 tensor of shape [hidden, expert_intermediate].
Contiguous BF16 tensor of shape [local_experts, expert_intermediate, hidden].
Contiguous BF16 tensor of shape [local_experts, expert_intermediate, hidden].
Contiguous BF16 tensor of shape [local_experts, hidden, expert_intermediate].
Prequantized gate, up, and down projections. Each is a
tuple (weight_fp8, weight_scale, weight_t_fp8, weight_t_scale)
in MoK’s opaque E4M3/E8M0 layouts. Must be non-None for MXFP8
and None for BF16.
Returns: torch.Tensor
Tuple of gradients for x, router weights, three routed weights,
Run MoK’s manual forward and retain its backward context.
Parameters:
Contiguous BF16 tensor of shape [tokens, hidden].
Contiguous FP32 tensor of shape [tokens, activated_experts].
Contiguous int64 tensor of shape [tokens, activated_experts].
Contiguous BF16 tensor of shape [expert_intermediate, hidden].
Contiguous BF16 tensor of shape [expert_intermediate, hidden].
Contiguous BF16 tensor of shape [hidden, expert_intermediate].
Contiguous BF16 tensor of shape [local_experts, expert_intermediate, hidden].
Contiguous BF16 tensor of shape [local_experts, expert_intermediate, hidden].
Contiguous BF16 tensor of shape [local_experts, hidden, expert_intermediate].
Returns: torch.Tensor
Tuple containing the BF16 output of shape [tokens, hidden], an opaque
Attach the runtime to the expert-parallel group.
Parameters:
One-dimensional device mesh named ep.
Global number of routed experts, divided evenly over the EP mesh.
Flatten a MoK schedule/context without retaining tensor references in metadata.
MoK’s dataclasses contain either tensors or tuples of tensors (the latter for
MXFP8 values). Saving the flattened tensors with ctx.save_for_backward
makes them visible to activation-checkpoint saved-tensor hooks. Storing the
original dataclass directly on ctx would keep every layer’s macro-sized
context alive until backward.
Parameters:
MoK schedule or forward-context dataclass.
Returns: tuple[tuple[torch.Tensor, ...], tuple[type, tuple[tuple[str, int], ...]]]
Flat tensors and tensor-free reconstruction metadata.
Import MoK after distributed setup has selected this process’s GPU.
Importing MoK loads its CUDA extension. Doing that at module-import time is
too early for torchrun workers: every local worker still has CUDA device 0
selected, leaving one stray CUDA context per worker on the first GPU. MoK
is only needed when expert parallelism is initialized, which happens after
initialize_distributed calls torch.cuda.set_device.
Import MoK’s weight-quantization ops after selecting this process’s GPU.
Return the rank-local tensor while preserving its autograd connection.
Parameters:
Tensor or DTensor of arbitrary shape. A DTensor must be sharded or replicated on the current rank with a materialized local value.
Returns: torch.Tensor
Tensor with the DTensor’s rank-local shape, or the original plain tensor.
Prequantize both MoK MXFP8 layouts in one kernel launch.
Parameters:
Contiguous BF16 tensor of shape [local_experts, output_features, input_features].
Returns: torch.Tensor
Normal E4M3 weight, its E8M0 block scales, transposed E4M3 weight,
Reconstruct a MoK dataclass from checkpoint-managed saved tensors.
Retain MXFP8 weights within a step and invalidate them after optimizer updates.
Parameters:
Pipeline model parts paired one-to-one with optimizers.
Optimizers whose post-step hooks own cache invalidation.