diff --git a/unsloth/kernels/moe/grouped_gemm/reference/layers/llama4_moe.py b/unsloth/kernels/moe/grouped_gemm/reference/layers/llama4_moe.py index 3d1afac9e3..6e833332e4 100644 --- a/unsloth/kernels/moe/grouped_gemm/reference/layers/llama4_moe.py +++ b/unsloth/kernels/moe/grouped_gemm/reference/layers/llama4_moe.py @@ -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