Added reference documentation for: unsloth/kernels/moe/grouped_gemm/reference/layers/qwen3_moe.py
This commit is contained in:
parent
962d444bb5
commit
0d91b0cad9
1 changed files with 137 additions and 1 deletions
|
|
@ -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
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue