Added reference documentation for: unsloth/kernels/moe/grouped_gemm/reference/layers/llama4_moe.py

This commit is contained in:
askmanu[bot] 2025-09-19 12:22:09 +00:00 committed by GitHub
commit 962d444bb5
No known key found for this signature in database
GPG key ID: B5690EEEBB952194

View file

@ -32,6 +32,24 @@ Reference implementation of Llama4 MoE block using triton grouped gemm.
@dataclass
class Llama4MoeResult:
"""Container for storing intermediate and final results from Llama4 MoE forward pass.
This dataclass holds all the intermediate tensors and computations from the MoE
forward pass, useful for debugging and analysis purposes.
Attributes:
token_counts_by_expert: Number of tokens assigned to each expert
gather_indices: Indices for gathering tokens in expert order
topk_weights: Top-k routing weights from the router
hidden_states_after_weight_merge: Hidden states after applying routing weights
first_gemm: Output of the first grouped GEMM operation
intermediate: Output after activation and multiplication
second_gemm: Output of the second grouped GEMM operation
hidden_states_unpermute: Hidden states after unpermuting from expert order
shared_expert_out: Output from the shared expert
final_out: Final output combining expert and shared expert outputs
router_logits: Raw logits from the router (optional)
"""
token_counts_by_expert: torch.Tensor
gather_indices: torch.Tensor
topk_weights: torch.Tensor
@ -46,6 +64,12 @@ class Llama4MoeResult:
class Llama4GroupedGemmTextMoe(Llama4TextMoe):
"""Llama4 MoE implementation using torch-native grouped GEMM operations.
This class extends the standard Llama4TextMoe with optimized grouped GEMM
operations for better performance. It permutes expert weights in-place for
optimal memory layout and supports optional router-shared expert overlap.
"""
EXPERT_WEIGHT_NAMES = ["experts.gate_up_proj", "experts.down_proj"]
def __init__(
@ -55,6 +79,14 @@ class Llama4GroupedGemmTextMoe(Llama4TextMoe):
verbose=False,
debug=False,
):
"""Initialize the Llama4GroupedGemmTextMoe layer.
Args:
config: Llama4 text configuration containing model parameters
overlap_router_shared: Whether to overlap router and shared expert computation
verbose: Whether to print detailed initialization information
debug: Whether to return detailed debug information in forward pass
"""
super().__init__(config)
self.overlap_router_shared = overlap_router_shared
self.verbose = verbose
@ -102,6 +134,17 @@ class Llama4GroupedGemmTextMoe(Llama4TextMoe):
@torch.no_grad
def copy_weights(self, other: Llama4TextMoe):
"""Copy weights from another Llama4TextMoe instance.
Copies all parameters from the source model, applying necessary permutations
for expert weights to match the expected layout.
Args:
other: Source Llama4TextMoe model to copy weights from
Returns:
Self for method chaining
"""
for name, param_to_copy in other.named_parameters():
if self.verbose:
print(f"Copying {name} with shape {param_to_copy.shape}")
@ -118,6 +161,17 @@ class Llama4GroupedGemmTextMoe(Llama4TextMoe):
return self
def check_weights(self, other: Llama4TextMoe):
"""Verify that weights match another Llama4TextMoe instance.
Compares all parameters with the reference model, applying necessary
permutations for expert weights before comparison.
Args:
other: Reference Llama4TextMoe model to compare against
Raises:
AssertionError: If any weights don't match or aren't contiguous
"""
for name, other_param in other.named_parameters():
if any(n in name for n in self.EXPERT_WEIGHT_NAMES):
other_param = other_param.permute(0, 2, 1)
@ -126,12 +180,34 @@ class Llama4GroupedGemmTextMoe(Llama4TextMoe):
assert param.is_contiguous(), f"{name} not contiguous!"
def act_and_mul(self, x: torch.Tensor) -> torch.Tensor:
"""Apply activation function to gate projection and multiply with up projection.
Splits the input tensor into gate and up projections, applies the activation
function to the gate projection, and multiplies with the up projection.
Args:
x: Input tensor with shape [..., 2 * expert_dim]
Returns:
Result of activation(gate_proj) * up_proj with shape [..., expert_dim]
"""
assert x.shape[-1] == 2 * self.experts.expert_dim
gate_proj = x[..., : self.experts.expert_dim]
up_proj = x[..., self.experts.expert_dim :]
return self.experts.act_fn(gate_proj) * up_proj
def run_router(self, hidden_states: torch.Tensor) -> torch.Tensor:
"""Run the router to compute routing weights and select experts.
Computes router logits, selects top-k experts, and applies sigmoid
activation to the routing weights.
Args:
hidden_states: Input hidden states tensor
Returns:
Tuple of (router_logits, routing_weights, selected_experts)
"""
# router_logits: (batch * sequence_length, n_experts)
hidden_states = hidden_states.view(-1, self.hidden_dim)
router_logits = self.router(hidden_states)
@ -146,6 +222,17 @@ class Llama4GroupedGemmTextMoe(Llama4TextMoe):
def get_token_counts_and_gather_indices(
self, selected_experts: torch.Tensor
) -> Tuple[torch.Tensor, torch.Tensor]:
"""Compute token counts per expert and gather indices for permutation.
Calculates how many tokens are assigned to each expert and generates
indices for gathering tokens in expert order.
Args:
selected_experts: Tensor of selected expert indices for each token
Returns:
Tuple of (token_counts_by_expert, gather_indices)
"""
token_counts_by_expert, gather_indices = get_routing_indices(
selected_experts, self.num_experts
)
@ -154,7 +241,17 @@ class Llama4GroupedGemmTextMoe(Llama4TextMoe):
return token_counts_by_expert, gather_indices
def forward(self, hidden_states: torch.Tensor) -> torch.Tensor:
""" """
"""Forward pass through the MoE layer.
Processes input through router, applies routing weights, performs grouped
GEMM operations through experts, and combines with shared expert output.
Args:
hidden_states: Input tensor with shape (batch_size, sequence_length, hidden_dim)
Returns:
Either Llama4MoeResult (if debug=True) or tuple of (final_output, routing_weights)
"""
batch_size, sequence_length, hidden_dim = hidden_states.shape
num_tokens = batch_size * sequence_length
total_tokens = num_tokens * self.top_k
@ -253,6 +350,12 @@ class Llama4GroupedGemmTextMoe(Llama4TextMoe):
class Llama4TritonTextMoe(Llama4GroupedGemmTextMoe):
"""Llama4 MoE implementation using Triton-optimized grouped GEMM operations.
This class extends Llama4GroupedGemmTextMoe to use Triton kernels for grouped
GEMM operations, providing better performance through kernel fusion and
optimized memory access patterns.
"""
def __init__(
self,
config: Llama4TextConfig,
@ -267,6 +370,21 @@ class Llama4TritonTextMoe(Llama4GroupedGemmTextMoe):
dX_only: bool = False,
verbose=False,
):
"""Initialize the Llama4TritonTextMoe layer.
Args:
config: Llama4 text configuration containing model parameters
overlap_router_shared: Whether to overlap router and shared expert computation
permute_x: Whether to permute input tensor (not supported for Llama4)
permute_y: Whether to permute output tensor
autotune: Whether to use kernel autotuning
kernel_config_fwd: Forward kernel configuration (required if autotune=False)
kernel_config_bwd_dW: Backward dW kernel configuration (required if autotune=False)
kernel_config_bwd_dX: Backward dX kernel configuration (required if autotune=False)
dW_only: Whether to compute only weight gradients
dX_only: Whether to compute only input gradients
verbose: Whether to print detailed information
"""
super().__init__(config, overlap_router_shared=overlap_router_shared)
assert not permute_x, (
"Llama4 triton grouped gemm does not support permute x due to pre-multiplication of router weights"
@ -288,6 +406,17 @@ class Llama4TritonTextMoe(Llama4GroupedGemmTextMoe):
@torch.no_grad
def copy_weights(self, other: Llama4TextMoe):
"""Copy weights from another Llama4TextMoe instance.
Copies all parameters from the source model, applying necessary permutations
for expert weights to match the expected layout.
Args:
other: Source Llama4TextMoe model to copy weights from
Returns:
Self for method chaining
"""
for name, param_to_copy in other.named_parameters():
if self.verbose:
print(f"Copying {name} with shape {param_to_copy.shape}")
@ -304,6 +433,17 @@ class Llama4TritonTextMoe(Llama4GroupedGemmTextMoe):
return self
def check_weights(self, other: Llama4TextMoe):
"""Verify that weights match another Llama4TextMoe instance.
Compares all parameters with the reference model, applying necessary
permutations for expert weights before comparison.
Args:
other: Reference Llama4TextMoe model to compare against
Raises:
AssertionError: If any weights don't match or aren't contiguous
"""
for name, other_param in other.named_parameters():
if any(n in name for n in self.EXPERT_WEIGHT_NAMES):
other_param = other_param.permute(0, 2, 1)
@ -312,12 +452,34 @@ class Llama4TritonTextMoe(Llama4GroupedGemmTextMoe):
assert param.is_contiguous(), f"{name} not contiguous!"
def act_and_mul(self, x: torch.Tensor) -> torch.Tensor:
"""Apply activation function to gate projection and multiply with up projection.
Splits the input tensor into gate and up projections, applies the activation
function to the gate projection, and multiplies with the up projection.
Args:
x: Input tensor with shape [..., 2 * expert_dim]
Returns:
Result of activation(gate_proj) * up_proj with shape [..., expert_dim]
"""
assert x.shape[-1] == 2 * self.experts.expert_dim
gate_proj = x[..., : self.experts.expert_dim]
up_proj = x[..., self.experts.expert_dim :]
return self.experts.act_fn(gate_proj) * up_proj
def run_router(self, hidden_states: torch.Tensor) -> torch.Tensor:
"""Run the router to compute routing weights and select experts.
Computes router logits, selects top-k experts, and applies sigmoid
activation to the routing weights.
Args:
hidden_states: Input hidden states tensor
Returns:
Tuple of (router_logits, routing_weights, selected_experts)
"""
# router_logits: (batch * sequence_length, n_experts)
hidden_states = hidden_states.view(-1, self.hidden_dim)
router_logits = self.router(hidden_states)
@ -332,6 +494,17 @@ class Llama4TritonTextMoe(Llama4GroupedGemmTextMoe):
def get_token_counts_and_gather_indices(
self, selected_experts: torch.Tensor
) -> Tuple[torch.Tensor, torch.Tensor]:
"""Compute token counts per expert and gather indices for permutation.
Calculates how many tokens are assigned to each expert and generates
indices for gathering tokens in expert order.
Args:
selected_experts: Tensor of selected expert indices for each token
Returns:
Tuple of (token_counts_by_expert, gather_indices)
"""
token_counts_by_expert, gather_indices = get_routing_indices(
selected_experts, self.num_experts
)
@ -340,7 +513,17 @@ class Llama4TritonTextMoe(Llama4GroupedGemmTextMoe):
return token_counts_by_expert, gather_indices
def forward(self, hidden_states: torch.Tensor) -> torch.Tensor:
""" """
"""Forward pass through the MoE layer using Triton grouped GEMM.
Processes input through router, applies routing weights, performs Triton-optimized
grouped GEMM operations through experts, and combines with shared expert output.
Args:
hidden_states: Input tensor with shape (batch_size, sequence_length, hidden_dim)
Returns:
Tuple of (final_output, routing_weights)
"""
batch_size, sequence_length, hidden_dim = hidden_states.shape
num_tokens = batch_size * sequence_length
total_tokens = num_tokens * self.top_k