Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
47 changes: 33 additions & 14 deletions src/infiniop/ops/mxfp4_common/cuda/fused_moe_mxfp4_kernel.cuh
Original file line number Diff line number Diff line change
Expand Up @@ -6,6 +6,7 @@

#include <cstddef>
#include <cstdint>
#include <limits>

namespace op::mxfp4_common::cuda {

Expand Down Expand Up @@ -35,8 +36,9 @@ __global__ void fusedMoeMxfp4W13Kernel(
size_t num_experts,
size_t hidden_size,
size_t intermediate_size,
infiniopFusedMoeActivation_t activation) {
const size_t block = blockIdx.x;
infiniopFusedMoeActivation_t activation,
size_t block_offset) {
const size_t block = block_offset + blockIdx.x;
const size_t route = block / intermediate_size;
const size_t i = block - route * intermediate_size;
if (route >= route_count || i >= intermediate_size) {
Expand Down Expand Up @@ -99,8 +101,9 @@ __global__ void fusedMoeMxfp4W2Kernel(
size_t topk,
size_t num_experts,
size_t hidden_size,
size_t intermediate_size) {
const size_t block = blockIdx.x;
size_t intermediate_size,
size_t block_offset) {
const size_t block = block_offset + blockIdx.x;
const size_t token = block / hidden_size;
const size_t h = block - token * hidden_size;
if (token >= num_tokens || h >= hidden_size) {
Expand Down Expand Up @@ -156,19 +159,35 @@ void launchFusedMoeMxfp4(
const op::fused_moe_mxfp4::FusedMoeMxfp4Info &info,
Stream stream) {
constexpr size_t block_size = 256;
// CUDA-compatible Hygon launches use a 32-bit total-thread range. Split
// large prefill grids instead of letting grid_size * block_size overflow.
constexpr size_t max_blocks_per_launch
= std::numeric_limits<uint32_t>::max() / block_size;
const size_t w13_grid = info.intermediate_size * info.routeCount();
fusedMoeMxfp4W13Kernel<<<w13_grid, block_size,
2 * block_size * sizeof(float), stream>>>(
activated, input, selected_experts, w13_packed, w13_scale,
info.routeCount(), info.topk, info.num_experts,
info.hidden_size, info.intermediate_size, info.activation);
for (size_t block_offset = 0; block_offset < w13_grid;
block_offset += max_blocks_per_launch) {
const size_t grid_size = (w13_grid - block_offset < max_blocks_per_launch)
? w13_grid - block_offset
: max_blocks_per_launch;
fusedMoeMxfp4W13Kernel<<<grid_size, block_size,
2 * block_size * sizeof(float), stream>>>(
activated, input, selected_experts, w13_packed, w13_scale,
info.routeCount(), info.topk, info.num_experts,
info.hidden_size, info.intermediate_size, info.activation, block_offset);
}

const size_t w2_grid = info.hidden_size * info.num_tokens;
fusedMoeMxfp4W2Kernel<<<w2_grid, block_size,
block_size * sizeof(float), stream>>>(
output, activated, selected_experts, routing_weights, w2_packed, w2_scale,
info.num_tokens, info.topk, info.num_experts,
info.hidden_size, info.intermediate_size);
for (size_t block_offset = 0; block_offset < w2_grid;
block_offset += max_blocks_per_launch) {
const size_t grid_size = (w2_grid - block_offset < max_blocks_per_launch)
? w2_grid - block_offset
: max_blocks_per_launch;
fusedMoeMxfp4W2Kernel<<<grid_size, block_size,
block_size * sizeof(float), stream>>>(
output, activated, selected_experts, routing_weights, w2_packed, w2_scale,
info.num_tokens, info.topk, info.num_experts,
info.hidden_size, info.intermediate_size, block_offset);
}
}

inline size_t fusedMoeMxfp4DtypeSize(infiniDtype_t dtype) {
Expand Down
124 changes: 114 additions & 10 deletions test/infinicore/ops/fused_moe_mxfp4.py
Original file line number Diff line number Diff line change
Expand Up @@ -3,7 +3,6 @@

sys.path.insert(0, os.path.join(os.path.dirname(__file__), ".."))

import infinicore
import torch
import torch.nn.functional as F
from framework import (
Expand All @@ -14,6 +13,7 @@
TestCase,
)

import infinicore

_DTYPES = [infinicore.float16, infinicore.bfloat16, infinicore.float32]
_TOLERANCE = {
Expand Down Expand Up @@ -80,24 +80,54 @@ def make_cases():
(7, 128, 64, 5, 2, 2, "prefill SiTU"),
]
cases = []
for T, H, I, E, topk, activation, description in configs:
input_data = torch.randn((T, H), generator=generator) * 0.2
ids = torch.randint(0, E, (T, topk), generator=generator, dtype=torch.int32)
if T > 1:
for (
num_tokens,
hidden_size,
intermediate_size,
num_experts,
topk,
activation,
description,
) in configs:
input_data = torch.randn((num_tokens, hidden_size), generator=generator) * 0.2
ids = torch.randint(
0,
num_experts,
(num_tokens, topk),
generator=generator,
dtype=torch.int32,
)
if num_tokens > 1:
ids[-1, -1] = -1
raw_routes = torch.rand((T, topk), generator=generator)
raw_routes = torch.rand((num_tokens, topk), generator=generator)
routing = raw_routes / raw_routes.sum(dim=-1, keepdim=True)
w13_packed = torch.randint(
0, 256, (E, 2 * I, H // 2), generator=generator, dtype=torch.uint8
0,
256,
(num_experts, 2 * intermediate_size, hidden_size // 2),
generator=generator,
dtype=torch.uint8,
)
w13_scale = torch.randint(
123, 129, (E, 2 * I, H // 32), generator=generator, dtype=torch.uint8
123,
129,
(num_experts, 2 * intermediate_size, hidden_size // 32),
generator=generator,
dtype=torch.uint8,
)
w2_packed = torch.randint(
0, 256, (E, H, I // 2), generator=generator, dtype=torch.uint8
0,
256,
(num_experts, hidden_size, intermediate_size // 2),
generator=generator,
dtype=torch.uint8,
)
w2_scale = torch.randint(
123, 129, (E, H, I // 32), generator=generator, dtype=torch.uint8
123,
129,
(num_experts, hidden_size, intermediate_size // 32),
generator=generator,
dtype=torch.uint8,
)
for dtype in _DTYPES:
tensors = [
Expand Down Expand Up @@ -130,6 +160,80 @@ def make_cases():
description=f"fused_moe_mxfp4 - {description} - dtype={dtype}",
)
)

# Kimi-K3 reaches this launch size during long prefills. CUDA-compatible
# backends limit a launch to a 32-bit total-thread range, so the W13 grid
# must be split once route_count * intermediate_size exceeds that limit.
num_tokens = 293
hidden_size = 64
intermediate_size = 3584
num_experts = 1
topk = 16
input_data = torch.randn((num_tokens, hidden_size), generator=generator) * 0.2
ids = torch.full((num_tokens, topk), -1, dtype=torch.int32)
ids[-1, -1] = 0
routing = torch.full((num_tokens, topk), 1.0 / topk)
tensors = [
(input_data, infinicore.bfloat16, "input"),
(ids, infinicore.int32, "selected_experts"),
(routing, infinicore.float32, "routing_weights"),
(
torch.full(
(num_experts, 2 * intermediate_size, hidden_size // 2),
0x22,
dtype=torch.uint8,
),
infinicore.uint8,
"w13_packed",
),
(
torch.full(
(num_experts, 2 * intermediate_size, hidden_size // 32),
127,
dtype=torch.uint8,
),
infinicore.uint8,
"w13_scale",
),
(
torch.full(
(num_experts, hidden_size, intermediate_size // 2),
0x22,
dtype=torch.uint8,
),
infinicore.uint8,
"w2_packed",
),
(
torch.full(
(num_experts, hidden_size, intermediate_size // 32),
127,
dtype=torch.uint8,
),
infinicore.uint8,
"w2_scale",
),
]
cases.append(
TestCase(
inputs=[
TensorSpec.from_tensor(
tuple(tensor.shape),
None,
tensor_dtype,
init_mode=TensorInitializer.MANUAL,
set_tensor=tensor,
name=name,
)
for tensor, tensor_dtype, name in tensors
],
kwargs={"activation": 2},
output_spec=None,
comparison_target=None,
tolerance=_TOLERANCE[infinicore.bfloat16],
description="fused_moe_mxfp4 - large-grid regression",
)
)
return cases


Expand Down
Loading