From 966affa466ba881962b936774376f14cbd50f9f1 Mon Sep 17 00:00:00 2001 From: "askmanu[bot]" <192355599+askmanu[bot]@users.noreply.github.com> Date: Fri, 19 Sep 2025 12:22:10 +0000 Subject: [PATCH 1/2] Added reference documentation for: unsloth/kernels/moe/grouped_gemm/reference/layers/qwen3_moe.py (cherry picked from commit 0d91b0cad97a292096bb4b8888067cd3191598fa) --- .../reference/layers/qwen3_moe.py | 138 +++++++++++++++++- 1 file changed, 137 insertions(+), 1 deletion(-) diff --git a/unsloth/kernels/moe/grouped_gemm/reference/layers/qwen3_moe.py b/unsloth/kernels/moe/grouped_gemm/reference/layers/qwen3_moe.py index 31c635ba37..1a0ea492b2 100644 --- a/unsloth/kernels/moe/grouped_gemm/reference/layers/qwen3_moe.py +++ b/unsloth/kernels/moe/grouped_gemm/reference/layers/qwen3_moe.py @@ -37,6 +37,18 @@ NOTE: This is NOT to be used for production as it contains many extra checks and @dataclass class GroupedGEMMResult: + """Container for storing intermediate and final results from grouped GEMM operations. + + Attributes: + token_counts_by_expert: Number of tokens assigned to each expert + gather_indices: Indices used for token permutation and unpermutation + topk_weights: Routing weights for the top-k selected experts + first_gemm: Output from the first grouped matrix multiplication + intermediate: Result after applying activation function and element-wise multiplication + second_gemm: Output from the second grouped matrix multiplication + hidden_states_unpermute: Hidden states after unpermutation from expert order to token order + hidden_states: Final output hidden states + """ token_counts_by_expert: torch.Tensor gather_indices: torch.Tensor topk_weights: torch.Tensor @@ -48,6 +60,12 @@ class GroupedGEMMResult: class Qwen3MoeGroupedGEMMBlock(torch.nn.Module): + """Reference implementation of Qwen3 Mixture of Experts block using grouped GEMM operations. + + This implementation uses torch-native operations and stores intermediate results for debugging. + It implements the MoE routing mechanism with top-k expert selection and grouped matrix multiplications. + """ + def __init__( self, config, @@ -55,6 +73,14 @@ class Qwen3MoeGroupedGEMMBlock(torch.nn.Module): gate_up_proj: torch.Tensor, down_proj: torch.Tensor, ): + """Initialize the Qwen3 MoE block with expert weights. + + Args: + config: Qwen3MoeConfig containing model configuration parameters + gate: Router gate weights for expert selection [num_experts, hidden_size] + gate_up_proj: Combined gate and up projection weights [num_experts, 2*moe_intermediate_size, hidden_size] + down_proj: Down projection weights [num_experts, hidden_size, moe_intermediate_size] + """ super().__init__() self.num_experts = config.num_experts self.top_k = config.num_experts_per_tok @@ -84,6 +110,17 @@ class Qwen3MoeGroupedGEMMBlock(torch.nn.Module): @staticmethod def extract_hf_weights(moe_block: Qwen3MoeSparseMoeBlock): + """Extract and reorganize weights from a HuggingFace Qwen3MoeSparseMoeBlock. + + Args: + moe_block: HuggingFace Qwen3MoeSparseMoeBlock instance + + Returns: + Tuple containing: + - gate: Router gate weights + - gate_up_proj: Combined gate and up projection weights + - down_proj: Down projection weights + """ config: Qwen3MoeConfig = moe_block.experts[0].config num_experts = config.num_experts @@ -105,11 +142,24 @@ class Qwen3MoeGroupedGEMMBlock(torch.nn.Module): @classmethod def from_hf(cls, moe_block: Qwen3MoeSparseMoeBlock): + """Create a Qwen3MoeGroupedGEMMBlock from a HuggingFace MoE block. + + Args: + moe_block: HuggingFace Qwen3MoeSparseMoeBlock instance + + Returns: + Qwen3MoeGroupedGEMMBlock instance with extracted weights + """ config: Qwen3MoeConfig = moe_block.experts[0].config gate, gate_up_proj, down_proj = cls.extract_hf_weights(moe_block) return cls(config, gate, gate_up_proj, down_proj) def check_weights(self, moe_block: Qwen3MoeSparseMoeBlock): + """Verify that the weights match those in the original HuggingFace MoE block. + + Args: + moe_block: HuggingFace Qwen3MoeSparseMoeBlock to compare against + """ for i in range(self.num_experts): assert self.gate_up_proj[i].equal( torch.cat( @@ -123,12 +173,31 @@ class Qwen3MoeGroupedGEMMBlock(torch.nn.Module): assert self.down_proj[i].equal(moe_block.experts[i].down_proj.weight.data) def act_and_mul(self, x: torch.Tensor) -> torch.Tensor: + """Apply activation function to gate projection and multiply with up projection. + + Args: + x: Input tensor with shape [..., 2 * moe_intermediate_size] + + Returns: + Result of activation(gate_proj) * up_proj with shape [..., moe_intermediate_size] + """ assert x.shape[-1] == 2 * self.moe_intermediate_size gate_proj = x[..., : self.moe_intermediate_size] up_proj = x[..., self.moe_intermediate_size :] return self.act_fn(gate_proj) * up_proj def run_router(self, hidden_states: torch.Tensor) -> torch.Tensor: + """Run the routing mechanism to select top-k experts for each token. + + Args: + hidden_states: Input hidden states [batch_size * seq_len, hidden_size] + + Returns: + Tuple containing: + - router_logits: Raw logits from the router + - routing_weights: Normalized weights for selected experts + - selected_experts: Indices of selected experts for each token + """ # router_logits: (batch * sequence_length, n_experts) router_logits = torch.nn.functional.linear(hidden_states, self.gate) @@ -146,6 +215,16 @@ class Qwen3MoeGroupedGEMMBlock(torch.nn.Module): 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. + + Args: + selected_experts: Indices of selected experts for each token + + Returns: + Tuple containing: + - token_counts_by_expert: Number of tokens assigned to each expert + - gather_indices: Indices for permuting tokens from token order to expert order + """ token_counts_by_expert, gather_indices = get_routing_indices( selected_experts, self.num_experts ) @@ -154,7 +233,16 @@ class Qwen3MoeGroupedGEMMBlock(torch.nn.Module): return token_counts_by_expert, gather_indices def forward(self, hidden_states: torch.Tensor) -> torch.Tensor: - """ """ + """Forward pass through the MoE block. + + Args: + hidden_states: Input tensor [batch_size, seq_len, hidden_size] + + Returns: + Tuple containing: + - GroupedGEMMResult: Container with all intermediate results + - router_logits: Raw routing logits for auxiliary loss computation + """ batch_size, sequence_length, hidden_dim = hidden_states.shape num_tokens = batch_size * sequence_length total_tokens = num_tokens * self.top_k @@ -214,6 +302,12 @@ class Qwen3MoeGroupedGEMMBlock(torch.nn.Module): class Qwen3MoeFusedGroupedGEMMBlock(Qwen3MoeGroupedGEMMBlock): + """Optimized Qwen3 MoE block using fused grouped GEMM kernels. + + This implementation uses Triton-based grouped GEMM kernels for improved performance + and supports various optimization options like permutation fusion and kernel tuning. + """ + def __init__( self, config: Qwen3MoeConfig, @@ -229,6 +323,22 @@ class Qwen3MoeFusedGroupedGEMMBlock(Qwen3MoeGroupedGEMMBlock): dW_only: bool = False, dX_only: bool = False, ): + """Initialize the fused grouped GEMM MoE block. + + Args: + config: Qwen3MoeConfig containing model configuration + gate: Router gate weights + gate_up_proj: Combined gate and up projection weights + down_proj: Down projection weights + permute_x: Whether to fuse input permutation in the first GEMM + permute_y: Whether to fuse output unpermutation in the second GEMM + autotune: Whether to automatically tune kernel configurations + kernel_config_fwd: Manual kernel configuration for forward pass + kernel_config_bwd_dW: Manual kernel configuration for weight gradients + kernel_config_bwd_dX: Manual kernel configuration for input gradients + dW_only: Whether to compute only weight gradients + dX_only: Whether to compute only input gradients + """ super().__init__(config, gate, gate_up_proj, down_proj) self.permute_x = permute_x self.permute_y = permute_y @@ -258,6 +368,22 @@ class Qwen3MoeFusedGroupedGEMMBlock(Qwen3MoeGroupedGEMMBlock): dW_only: bool = False, dX_only: bool = False, ): + """Create a fused grouped GEMM MoE block from a HuggingFace MoE block. + + Args: + moe_block: HuggingFace Qwen3MoeSparseMoeBlock instance + permute_x: Whether to fuse input permutation in the first GEMM + permute_y: Whether to fuse output unpermutation in the second GEMM + autotune: Whether to automatically tune kernel configurations + kernel_config_fwd: Manual kernel configuration for forward pass + kernel_config_bwd_dW: Manual kernel configuration for weight gradients + kernel_config_bwd_dX: Manual kernel configuration for input gradients + dW_only: Whether to compute only weight gradients + dX_only: Whether to compute only input gradients + + Returns: + Qwen3MoeFusedGroupedGEMMBlock instance with extracted weights and configurations + """ config: Qwen3MoeConfig = moe_block.experts[0].config gate, gate_up_proj, down_proj = Qwen3MoeGroupedGEMMBlock.extract_hf_weights( moe_block @@ -278,6 +404,16 @@ class Qwen3MoeFusedGroupedGEMMBlock(Qwen3MoeGroupedGEMMBlock): ) def forward(self, hidden_states: torch.Tensor) -> torch.Tensor: + """Forward pass using fused grouped GEMM kernels. + + Args: + hidden_states: Input tensor [batch_size, seq_len, hidden_size] + + Returns: + Tuple containing: + - hidden_states: Output tensor [batch_size, seq_len, hidden_size] + - router_logits: Raw routing logits for auxiliary loss computation + """ batch_size, sequence_length, hidden_dim = hidden_states.shape num_tokens = batch_size * sequence_length total_tokens = num_tokens * self.top_k From e558b71df96d7f6bde2703a9da3b94a2cc1fbefa Mon Sep 17 00:00:00 2001 From: "pre-commit-ci[bot]" <66853113+pre-commit-ci[bot]@users.noreply.github.com> Date: Thu, 12 Mar 2026 07:51:42 +0000 Subject: [PATCH 2/2] [pre-commit.ci] auto fixes from pre-commit.com hooks for more information, see https://pre-commit.ci --- .../reference/layers/qwen3_moe.py | 49 ++++++++++--------- 1 file changed, 25 insertions(+), 24 deletions(-) diff --git a/unsloth/kernels/moe/grouped_gemm/reference/layers/qwen3_moe.py b/unsloth/kernels/moe/grouped_gemm/reference/layers/qwen3_moe.py index 1a0ea492b2..62e9b25bcf 100644 --- a/unsloth/kernels/moe/grouped_gemm/reference/layers/qwen3_moe.py +++ b/unsloth/kernels/moe/grouped_gemm/reference/layers/qwen3_moe.py @@ -38,7 +38,7 @@ NOTE: This is NOT to be used for production as it contains many extra checks and @dataclass class GroupedGEMMResult: """Container for storing intermediate and final results from grouped GEMM operations. - + Attributes: token_counts_by_expert: Number of tokens assigned to each expert gather_indices: Indices used for token permutation and unpermutation @@ -49,6 +49,7 @@ class GroupedGEMMResult: hidden_states_unpermute: Hidden states after unpermutation from expert order to token order hidden_states: Final output hidden states """ + token_counts_by_expert: torch.Tensor gather_indices: torch.Tensor topk_weights: torch.Tensor @@ -61,11 +62,11 @@ class GroupedGEMMResult: class Qwen3MoeGroupedGEMMBlock(torch.nn.Module): """Reference implementation of Qwen3 Mixture of Experts block using grouped GEMM operations. - + This implementation uses torch-native operations and stores intermediate results for debugging. It implements the MoE routing mechanism with top-k expert selection and grouped matrix multiplications. """ - + def __init__( self, config, @@ -74,7 +75,7 @@ class Qwen3MoeGroupedGEMMBlock(torch.nn.Module): down_proj: torch.Tensor, ): """Initialize the Qwen3 MoE block with expert weights. - + Args: config: Qwen3MoeConfig containing model configuration parameters gate: Router gate weights for expert selection [num_experts, hidden_size] @@ -111,10 +112,10 @@ class Qwen3MoeGroupedGEMMBlock(torch.nn.Module): @staticmethod def extract_hf_weights(moe_block: Qwen3MoeSparseMoeBlock): """Extract and reorganize weights from a HuggingFace Qwen3MoeSparseMoeBlock. - + Args: moe_block: HuggingFace Qwen3MoeSparseMoeBlock instance - + Returns: Tuple containing: - gate: Router gate weights @@ -143,10 +144,10 @@ class Qwen3MoeGroupedGEMMBlock(torch.nn.Module): @classmethod def from_hf(cls, moe_block: Qwen3MoeSparseMoeBlock): """Create a Qwen3MoeGroupedGEMMBlock from a HuggingFace MoE block. - + Args: moe_block: HuggingFace Qwen3MoeSparseMoeBlock instance - + Returns: Qwen3MoeGroupedGEMMBlock instance with extracted weights """ @@ -156,7 +157,7 @@ class Qwen3MoeGroupedGEMMBlock(torch.nn.Module): def check_weights(self, moe_block: Qwen3MoeSparseMoeBlock): """Verify that the weights match those in the original HuggingFace MoE block. - + Args: moe_block: HuggingFace Qwen3MoeSparseMoeBlock to compare against """ @@ -174,10 +175,10 @@ class Qwen3MoeGroupedGEMMBlock(torch.nn.Module): def act_and_mul(self, x: torch.Tensor) -> torch.Tensor: """Apply activation function to gate projection and multiply with up projection. - + Args: x: Input tensor with shape [..., 2 * moe_intermediate_size] - + Returns: Result of activation(gate_proj) * up_proj with shape [..., moe_intermediate_size] """ @@ -188,10 +189,10 @@ class Qwen3MoeGroupedGEMMBlock(torch.nn.Module): def run_router(self, hidden_states: torch.Tensor) -> torch.Tensor: """Run the routing mechanism to select top-k experts for each token. - + Args: hidden_states: Input hidden states [batch_size * seq_len, hidden_size] - + Returns: Tuple containing: - router_logits: Raw logits from the router @@ -216,10 +217,10 @@ class Qwen3MoeGroupedGEMMBlock(torch.nn.Module): self, selected_experts: torch.Tensor ) -> Tuple[torch.Tensor, torch.Tensor]: """Compute token counts per expert and gather indices for permutation. - + Args: selected_experts: Indices of selected experts for each token - + Returns: Tuple containing: - token_counts_by_expert: Number of tokens assigned to each expert @@ -234,10 +235,10 @@ class Qwen3MoeGroupedGEMMBlock(torch.nn.Module): def forward(self, hidden_states: torch.Tensor) -> torch.Tensor: """Forward pass through the MoE block. - + Args: hidden_states: Input tensor [batch_size, seq_len, hidden_size] - + Returns: Tuple containing: - GroupedGEMMResult: Container with all intermediate results @@ -303,11 +304,11 @@ class Qwen3MoeGroupedGEMMBlock(torch.nn.Module): class Qwen3MoeFusedGroupedGEMMBlock(Qwen3MoeGroupedGEMMBlock): """Optimized Qwen3 MoE block using fused grouped GEMM kernels. - + This implementation uses Triton-based grouped GEMM kernels for improved performance and supports various optimization options like permutation fusion and kernel tuning. """ - + def __init__( self, config: Qwen3MoeConfig, @@ -324,7 +325,7 @@ class Qwen3MoeFusedGroupedGEMMBlock(Qwen3MoeGroupedGEMMBlock): dX_only: bool = False, ): """Initialize the fused grouped GEMM MoE block. - + Args: config: Qwen3MoeConfig containing model configuration gate: Router gate weights @@ -369,7 +370,7 @@ class Qwen3MoeFusedGroupedGEMMBlock(Qwen3MoeGroupedGEMMBlock): dX_only: bool = False, ): """Create a fused grouped GEMM MoE block from a HuggingFace MoE block. - + Args: moe_block: HuggingFace Qwen3MoeSparseMoeBlock instance permute_x: Whether to fuse input permutation in the first GEMM @@ -380,7 +381,7 @@ class Qwen3MoeFusedGroupedGEMMBlock(Qwen3MoeGroupedGEMMBlock): kernel_config_bwd_dX: Manual kernel configuration for input gradients dW_only: Whether to compute only weight gradients dX_only: Whether to compute only input gradients - + Returns: Qwen3MoeFusedGroupedGEMMBlock instance with extracted weights and configurations """ @@ -405,10 +406,10 @@ class Qwen3MoeFusedGroupedGEMMBlock(Qwen3MoeGroupedGEMMBlock): def forward(self, hidden_states: torch.Tensor) -> torch.Tensor: """Forward pass using fused grouped GEMM kernels. - + Args: hidden_states: Input tensor [batch_size, seq_len, hidden_size] - + Returns: Tuple containing: - hidden_states: Output tensor [batch_size, seq_len, hidden_size]