Compare commits
2 commits
main
...
dh/recover
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
e558b71df9 | ||
|
|
966affa466 |
1 changed files with 138 additions and 1 deletions
|
|
@ -37,6 +37,19 @@ NOTE: This is NOT to be used for production as it contains many extra checks and
|
||||||
|
|
||||||
@dataclass
|
@dataclass
|
||||||
class GroupedGEMMResult:
|
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
|
token_counts_by_expert: torch.Tensor
|
||||||
gather_indices: torch.Tensor
|
gather_indices: torch.Tensor
|
||||||
topk_weights: torch.Tensor
|
topk_weights: torch.Tensor
|
||||||
|
|
@ -48,6 +61,12 @@ class GroupedGEMMResult:
|
||||||
|
|
||||||
|
|
||||||
class Qwen3MoeGroupedGEMMBlock(torch.nn.Module):
|
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__(
|
def __init__(
|
||||||
self,
|
self,
|
||||||
config,
|
config,
|
||||||
|
|
@ -55,6 +74,14 @@ class Qwen3MoeGroupedGEMMBlock(torch.nn.Module):
|
||||||
gate_up_proj: torch.Tensor,
|
gate_up_proj: torch.Tensor,
|
||||||
down_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__()
|
super().__init__()
|
||||||
self.num_experts = config.num_experts
|
self.num_experts = config.num_experts
|
||||||
self.top_k = config.num_experts_per_tok
|
self.top_k = config.num_experts_per_tok
|
||||||
|
|
@ -84,6 +111,17 @@ class Qwen3MoeGroupedGEMMBlock(torch.nn.Module):
|
||||||
|
|
||||||
@staticmethod
|
@staticmethod
|
||||||
def extract_hf_weights(moe_block: Qwen3MoeSparseMoeBlock):
|
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
|
config: Qwen3MoeConfig = moe_block.experts[0].config
|
||||||
num_experts = config.num_experts
|
num_experts = config.num_experts
|
||||||
|
|
||||||
|
|
@ -105,11 +143,24 @@ class Qwen3MoeGroupedGEMMBlock(torch.nn.Module):
|
||||||
|
|
||||||
@classmethod
|
@classmethod
|
||||||
def from_hf(cls, moe_block: Qwen3MoeSparseMoeBlock):
|
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
|
config: Qwen3MoeConfig = moe_block.experts[0].config
|
||||||
gate, gate_up_proj, down_proj = cls.extract_hf_weights(moe_block)
|
gate, gate_up_proj, down_proj = cls.extract_hf_weights(moe_block)
|
||||||
return cls(config, gate, gate_up_proj, down_proj)
|
return cls(config, gate, gate_up_proj, down_proj)
|
||||||
|
|
||||||
def check_weights(self, moe_block: Qwen3MoeSparseMoeBlock):
|
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):
|
for i in range(self.num_experts):
|
||||||
assert self.gate_up_proj[i].equal(
|
assert self.gate_up_proj[i].equal(
|
||||||
torch.cat(
|
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)
|
assert self.down_proj[i].equal(moe_block.experts[i].down_proj.weight.data)
|
||||||
|
|
||||||
def act_and_mul(self, x: torch.Tensor) -> torch.Tensor:
|
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
|
assert x.shape[-1] == 2 * self.moe_intermediate_size
|
||||||
gate_proj = x[..., : self.moe_intermediate_size]
|
gate_proj = x[..., : self.moe_intermediate_size]
|
||||||
up_proj = x[..., self.moe_intermediate_size :]
|
up_proj = x[..., self.moe_intermediate_size :]
|
||||||
return self.act_fn(gate_proj) * up_proj
|
return self.act_fn(gate_proj) * up_proj
|
||||||
|
|
||||||
def run_router(self, hidden_states: torch.Tensor) -> torch.Tensor:
|
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: (batch * sequence_length, n_experts)
|
||||||
router_logits = torch.nn.functional.linear(hidden_states, self.gate)
|
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(
|
def get_token_counts_and_gather_indices(
|
||||||
self, selected_experts: torch.Tensor
|
self, selected_experts: torch.Tensor
|
||||||
) -> Tuple[torch.Tensor, 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(
|
token_counts_by_expert, gather_indices = get_routing_indices(
|
||||||
selected_experts, self.num_experts
|
selected_experts, self.num_experts
|
||||||
)
|
)
|
||||||
|
|
@ -154,7 +234,16 @@ class Qwen3MoeGroupedGEMMBlock(torch.nn.Module):
|
||||||
return token_counts_by_expert, gather_indices
|
return token_counts_by_expert, gather_indices
|
||||||
|
|
||||||
def forward(self, hidden_states: torch.Tensor) -> torch.Tensor:
|
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
|
batch_size, sequence_length, hidden_dim = hidden_states.shape
|
||||||
num_tokens = batch_size * sequence_length
|
num_tokens = batch_size * sequence_length
|
||||||
total_tokens = num_tokens * self.top_k
|
total_tokens = num_tokens * self.top_k
|
||||||
|
|
@ -214,6 +303,12 @@ class Qwen3MoeGroupedGEMMBlock(torch.nn.Module):
|
||||||
|
|
||||||
|
|
||||||
class Qwen3MoeFusedGroupedGEMMBlock(Qwen3MoeGroupedGEMMBlock):
|
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__(
|
def __init__(
|
||||||
self,
|
self,
|
||||||
config: Qwen3MoeConfig,
|
config: Qwen3MoeConfig,
|
||||||
|
|
@ -229,6 +324,22 @@ class Qwen3MoeFusedGroupedGEMMBlock(Qwen3MoeGroupedGEMMBlock):
|
||||||
dW_only: bool = False,
|
dW_only: bool = False,
|
||||||
dX_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)
|
super().__init__(config, gate, gate_up_proj, down_proj)
|
||||||
self.permute_x = permute_x
|
self.permute_x = permute_x
|
||||||
self.permute_y = permute_y
|
self.permute_y = permute_y
|
||||||
|
|
@ -258,6 +369,22 @@ class Qwen3MoeFusedGroupedGEMMBlock(Qwen3MoeGroupedGEMMBlock):
|
||||||
dW_only: bool = False,
|
dW_only: bool = False,
|
||||||
dX_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
|
config: Qwen3MoeConfig = moe_block.experts[0].config
|
||||||
gate, gate_up_proj, down_proj = Qwen3MoeGroupedGEMMBlock.extract_hf_weights(
|
gate, gate_up_proj, down_proj = Qwen3MoeGroupedGEMMBlock.extract_hf_weights(
|
||||||
moe_block
|
moe_block
|
||||||
|
|
@ -278,6 +405,16 @@ class Qwen3MoeFusedGroupedGEMMBlock(Qwen3MoeGroupedGEMMBlock):
|
||||||
)
|
)
|
||||||
|
|
||||||
def forward(self, hidden_states: torch.Tensor) -> torch.Tensor:
|
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
|
batch_size, sequence_length, hidden_dim = hidden_states.shape
|
||||||
num_tokens = batch_size * sequence_length
|
num_tokens = batch_size * sequence_length
|
||||||
total_tokens = num_tokens * self.top_k
|
total_tokens = num_tokens * self.top_k
|
||||||
|
|
|
||||||
Loading…
Add table
Add a link
Reference in a new issue