Compare commits

...
Sign in to create a new pull request.

2 commits

Author SHA1 Message Date
pre-commit-ci[bot]
e558b71df9 [pre-commit.ci] auto fixes from pre-commit.com hooks
for more information, see https://pre-commit.ci
2026-03-12 07:51:43 +00:00
askmanu[bot]
966affa466 Added reference documentation for: unsloth/kernels/moe/grouped_gemm/reference/layers/qwen3_moe.py
(cherry picked from commit 0d91b0cad9)
2026-03-12 07:48:58 +00:00

View file

@ -37,6 +37,19 @@ 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 +61,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 +74,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 +111,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 +143,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 +174,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 +216,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 +234,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 +303,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 +324,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 +369,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 +405,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