diff --git a/src/infiniop/ops/mxfp4_common/cuda/fused_moe_mxfp4_kernel.cuh b/src/infiniop/ops/mxfp4_common/cuda/fused_moe_mxfp4_kernel.cuh index 616521859..8d60a68aa 100644 --- a/src/infiniop/ops/mxfp4_common/cuda/fused_moe_mxfp4_kernel.cuh +++ b/src/infiniop/ops/mxfp4_common/cuda/fused_moe_mxfp4_kernel.cuh @@ -6,6 +6,7 @@ #include #include +#include namespace op::mxfp4_common::cuda { @@ -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) { @@ -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) { @@ -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::max() / block_size; const size_t w13_grid = info.intermediate_size * info.routeCount(); - fusedMoeMxfp4W13Kernel<<>>( - 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<<>>( + 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<<>>( - 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<<>>( + 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) { diff --git a/test/infinicore/ops/fused_moe_mxfp4.py b/test/infinicore/ops/fused_moe_mxfp4.py index 82bd7a66e..2aa8cdae9 100644 --- a/test/infinicore/ops/fused_moe_mxfp4.py +++ b/test/infinicore/ops/fused_moe_mxfp4.py @@ -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 ( @@ -14,6 +13,7 @@ TestCase, ) +import infinicore _DTYPES = [infinicore.float16, infinicore.bfloat16, infinicore.float32] _TOLERANCE = { @@ -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 = [ @@ -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