Download kernels/moe_kernel.py from Snapkitty/snapkitty-transformer: direct link, hf CLI and curl.
- Browser
- Download file 4.34 kB
-
https://huggingface.co/Snapkitty/snapkitty-transformer/resolve/main/kernels/moe_kernel.py
- Command line
-
hf download hf://Snapkitty/snapkitty-transformer/kernels/moe_kernel.py
-
curl -L -o moe_kernel.py https://huggingface.co/Snapkitty/snapkitty-transformer/resolve/main/kernels/moe_kernel.py
4.34 kB
| import triton | |
| import triton.language as tl | |
| import torch | |
| def _moe_kernel( | |
| X, W_gate, W_up, W_down, W_gate_scale, W_up_scale, W_down_scale, | |
| top_k_indices, top_k_weights, | |
| O, | |
| stride_xb, stride_xd, | |
| stride_og, stride_od, | |
| num_experts: tl.constexpr, | |
| top_k: tl.constexpr, | |
| dim: tl.constexpr, | |
| intermediate_dim: tl.constexpr, | |
| use_mxfp4: tl.constexpr, | |
| BLOCK_SIZE: tl.constexpr, | |
| BLOCK_INTER: tl.constexpr, | |
| ): | |
| pid = tl.program_id(0) | |
| offs_d = tl.arange(0, BLOCK_SIZE) | |
| offs_inter = tl.arange(0, BLOCK_INTER) | |
| x = tl.load(X + pid * stride_xb + offs_d, mask=offs_d < dim, other=0.0).to(tl.float32) | |
| acc = tl.zeros([BLOCK_SIZE], dtype=tl.float32) | |
| for k in range(top_k): | |
| expert_idx = tl.load(top_k_indices + pid * top_k + k) | |
| expert_weight = tl.load(top_k_weights + pid * top_k + k) | |
| w_gate_ptrs = W_gate + expert_idx * dim * dim + offs_inter[:, None] * dim + offs_d[None, :] | |
| w_gate = tl.load(w_gate_ptrs, mask=(offs_inter[:, None] < intermediate_dim) & (offs_d[None, :] < dim), other=0.0) | |
| if use_mxfp4: | |
| scale = tl.load(W_gate_scale + expert_idx) | |
| w_gate = w_gate.to(tl.float32) * scale | |
| gate_out = tl.dot(x, tl.trans(w_gate)) | |
| gate_out = tl.sigmoid(gate_out) * gate_out | |
| w_up_ptrs = W_up + expert_idx * dim * dim + offs_inter[:, None] * dim + offs_d[None, :] | |
| w_up = tl.load(w_up_ptrs, mask=(offs_inter[:, None] < intermediate_dim) & (offs_d[None, :] < dim), other=0.0) | |
| if use_mxfp4: | |
| scale = tl.load(W_up_scale + expert_idx) | |
| w_up = w_up.to(tl.float32) * scale | |
| up_out = tl.dot(x, tl.trans(w_up)) | |
| hidden = gate_out * up_out | |
| w_down_ptrs = W_down + expert_idx * dim * dim + offs_d[:, None] * offs_inter[None, :] | |
| w_down = tl.load(w_down_ptrs, mask=(offs_d[:, None] < dim) & (offs_inter[None, :] < intermediate_dim), other=0.0) | |
| if use_mxfp4: | |
| scale = tl.load(W_down_scale + expert_idx) | |
| w_down = w_down.to(tl.float32) * scale | |
| expert_out = tl.dot(hidden, tl.trans(w_down)) | |
| acc += expert_out * expert_weight | |
| tl.store(O + pid * stride_og + offs_d, acc, mask=offs_d < dim) | |
| def fused_moe_triton( | |
| X, W_gate, W_up, W_down, | |
| top_k_indices, top_k_weights, | |
| W_gate_scale=None, W_up_scale=None, W_down_scale=None, | |
| use_mxfp4=False | |
| ): | |
| batch_size, dim = X.shape | |
| num_experts = W_gate.shape[0] | |
| intermediate_dim = W_gate.shape[1] | |
| top_k = top_k_indices.shape[1] | |
| BLOCK_SIZE = triton.next_power_of_2(dim) | |
| BLOCK_INTER = triton.next_power_of_2(intermediate_dim) | |
| O = torch.zeros_like(X) | |
| grid = (batch_size,) | |
| _moe_kernel[grid]( | |
| X, W_gate, W_up, W_down, | |
| W_gate_scale if W_gate_scale is not None else X, | |
| W_up_scale if W_up_scale is not None else X, | |
| W_down_scale if W_down_scale is not None else X, | |
| top_k_indices, top_k_weights, | |
| O, | |
| X.stride(0), X.stride(1), | |
| O.stride(0), O.stride(1), | |
| num_experts=num_experts, | |
| top_k=top_k, | |
| dim=dim, | |
| intermediate_dim=intermediate_dim, | |
| use_mxfp4=use_mxfp4, | |
| BLOCK_SIZE=BLOCK_SIZE, | |
| BLOCK_INTER=BLOCK_INTER, | |
| ) | |
| return O | |
| def test_moe(): | |
| torch.manual_seed(42) | |
| batch, dim, intermediate, num_experts, top_k = 32, 512, 1024, 8, 4 | |
| X = torch.randn(batch, dim, dtype=torch.float16, device='cuda') | |
| W_gate = torch.randn(num_experts, dim, intermediate, dtype=torch.float16, device='cuda') | |
| W_up = torch.randn(num_experts, dim, intermediate, dtype=torch.float16, device='cuda') | |
| W_down = torch.randn(num_experts, intermediate, dim, dtype=torch.float16, device='cuda') | |
| gate_logits = torch.randn(batch, num_experts, device='cuda') | |
| top_k_weights, top_k_indices = torch.topk(F.softmax(gate_logits, dim=-1), top_k) | |
| top_k_weights = top_k_weights / top_k_weights.sum(dim=-1, keepdim=True) | |
| O = fused_moe_triton(X, W_gate, W_up, W_down, top_k_indices, top_k_weights) | |
| assert O.shape == X.shape | |
| print(f"MoE output shape: {O.shape}") | |
| print(f"MoE test passed!") | |
| if __name__ == '__main__': | |
| test_moe() |