From 057ae2101ea0688e1820478c39f2daca73423e66 Mon Sep 17 00:00:00 2001 From: biondizzle Date: Mon, 1 Jun 2026 10:28:01 +0000 Subject: [PATCH] CRITICAL FIX: Move tiled_mma creation and _setup_attributes OUTSIDE @cute.jit The _setup_attributes() calls cute.size(tiled_mma.shape_mnk, mode=[2]) which requires host-side execution. Inside @cute.jit, tiled_mma.shape_mnk returns MLIR values that can't be unpacked by cute.size(). This follows the fused_swiglu.py pattern exactly: setup on host side, then pass everything to the kernel. Removed @cute.jit wrapper entirely in favor of direct kernel launch (same as fused_swiglu). --- .../router/nvfp4_fused_router_kernel.py | 100 +++++++++--------- 1 file changed, 50 insertions(+), 50 deletions(-) diff --git a/dsv4/kernels/router/nvfp4_fused_router_kernel.py b/dsv4/kernels/router/nvfp4_fused_router_kernel.py index 3445bb82..10c737f2 100644 --- a/dsv4/kernels/router/nvfp4_fused_router_kernel.py +++ b/dsv4/kernels/router/nvfp4_fused_router_kernel.py @@ -287,63 +287,63 @@ class Nvfp4FusedRouterKernel: num_N_tiles = (N + cta_n - 1) // cta_n grid = (num_M_tiles * num_N_tiles, 1, 1) - @cute.jit - def _compiled_fn(mat_a, mat_b, scale_a, scale_b, mat_c): - tiled_mma = self._create_tiled_mma(a_dtype, a_major_mode, b_major_mode, sf_dtype) - tiled_mma_sfb = self._create_tiled_mma_sfb(a_dtype, a_major_mode, b_major_mode, sf_dtype) - self._setup_attributes(tiled_mma, tiled_mma_sfb, a_dtype, b_dtype, sf_dtype, c_dtype, c_layout) + # Setup tiled MMA and attributes on HOST side (outside JIT) + # Same pattern as fused_swiglu.py __call__ + # _setup_attributes calls cute.size(tiled_mma.shape_mnk, mode=[2]) + # which requires host-side execution (not inside @cute.jit) + tiled_mma = self._create_tiled_mma(a_dtype, a_major_mode, b_major_mode, sf_dtype) + tiled_mma_sfb = self._create_tiled_mma_sfb(a_dtype, a_major_mode, b_major_mode, sf_dtype) + self._setup_attributes(tiled_mma, tiled_mma_sfb, a_dtype, b_dtype, sf_dtype, c_dtype, c_layout) - # TMA atoms for A, B, SFA, SFB - a_op = sm100_utils.cluster_shape_to_tma_atom_A(self.cluster_shape_mn, tiled_mma.thr_id) - a_smem_layout = cute.slice_(self.a_smem_layout_staged, (None, None, None, 0)) - tma_atom_a, tma_tensor_a = cute.nvgpu.make_tiled_tma_atom_A( - a_op, mat_a, a_smem_layout, self.mma_tiler, tiled_mma, self.cluster_layout_vmnk.shape) + # TMA atoms (host side, same as fused_swiglu) + a_op = sm100_utils.cluster_shape_to_tma_atom_A(self.cluster_shape_mn, tiled_mma.thr_id) + a_smem_layout = cute.slice_(self.a_smem_layout_staged, (None, None, None, 0)) + tma_atom_a, tma_tensor_a = cute.nvgpu.make_tiled_tma_atom_A( + a_op, mat_a, a_smem_layout, self.mma_tiler, tiled_mma, self.cluster_layout_vmnk.shape) - b_op = sm100_utils.cluster_shape_to_tma_atom_B(self.cluster_shape_mn, tiled_mma.thr_id) - b_smem_layout = cute.slice_(self.b_smem_layout_staged, (None, None, None, 0)) - tma_atom_b, tma_tensor_b = cute.nvgpu.make_tiled_tma_atom_B( - b_op, mat_b, b_smem_layout, self.mma_tiler, tiled_mma, self.cluster_layout_vmnk.shape) + b_op = sm100_utils.cluster_shape_to_tma_atom_B(self.cluster_shape_mn, tiled_mma.thr_id) + b_smem_layout = cute.slice_(self.b_smem_layout_staged, (None, None, None, 0)) + tma_atom_b, tma_tensor_b = cute.nvgpu.make_tiled_tma_atom_B( + b_op, mat_b, b_smem_layout, self.mma_tiler, tiled_mma, self.cluster_layout_vmnk.shape) - sfa_op = sm100_utils.cluster_shape_to_tma_atom_A(self.cluster_shape_mn, tiled_mma.thr_id) - sfa_smem_layout = cute.slice_(self.sfa_smem_layout_staged, (None, None, None, 0)) - tma_atom_sfa, tma_tensor_sfa = cute.nvgpu.make_tiled_tma_atom_A( - sfa_op, scale_a, sfa_smem_layout, self.mma_tiler, tiled_mma, self.cluster_layout_vmnk.shape, - internal_type=cutlass.Uint64) + sfa_op = sm100_utils.cluster_shape_to_tma_atom_A(self.cluster_shape_mn, tiled_mma.thr_id) + sfa_smem_layout = cute.slice_(self.sfa_smem_layout_staged, (None, None, None, 0)) + tma_atom_sfa, tma_tensor_sfa = cute.nvgpu.make_tiled_tma_atom_A( + sfa_op, scale_a, sfa_smem_layout, self.mma_tiler, tiled_mma, self.cluster_layout_vmnk.shape, + internal_type=cutlass.Uint64) - sfb_op = sm100_utils.cluster_shape_to_tma_atom_SFB(self.cluster_shape_mn, tiled_mma.thr_id) - sfb_smem_layout = cute.slice_(self.sfb_smem_layout_staged, (None, None, None, 0)) - tma_atom_sfb, tma_tensor_sfb = cute.nvgpu.make_tiled_tma_atom_B( - sfb_op, scale_b, sfb_smem_layout, self.mma_tiler_sfb, tiled_mma_sfb, - self.cluster_layout_sfb_vmnk.shape, internal_type=cutlass.Uint64) + sfb_op = sm100_utils.cluster_shape_to_tma_atom_SFB(self.cluster_shape_mn, tiled_mma.thr_id) + sfb_smem_layout = cute.slice_(self.sfb_smem_layout_staged, (None, None, None, 0)) + tma_atom_sfb, tma_tensor_sfb = cute.nvgpu.make_tiled_tma_atom_B( + sfb_op, scale_b, sfb_smem_layout, self.mma_tiler_sfb, tiled_mma_sfb, + self.cluster_layout_sfb_vmnk.shape, internal_type=cutlass.Uint64) - # TMA store for C (activated scores) - epi_smem_layout = cute.slice_(self.c_smem_layout_staged, (None, None, 0)) - tma_atom_c, tma_tensor_c = cpasync.make_tiled_tma_atom( - cpasync.CopyBulkTensorTileS2GOp(), mat_c, epi_smem_layout, self.epi_tile) + epi_smem_layout = cute.slice_(self.c_smem_layout_staged, (None, None, 0)) + tma_atom_c, tma_tensor_c = cpasync.make_tiled_tma_atom( + cpasync.CopyBulkTensorTileS2GOp(), mat_c, epi_smem_layout, self.epi_tile) - tile_sched_params = utils.PersistentTileSchedulerParams( - (cutlass.Int32(num_M_tiles), cutlass.Int32(num_N_tiles), cutlass.Int32(1)), - (1, 1, 1)) + tile_sched_params = utils.PersistentTileSchedulerParams( + (cutlass.Int32(num_M_tiles), cutlass.Int32(num_N_tiles), cutlass.Int32(1)), + (1, 1, 1)) - self._kernel( - tiled_mma, tiled_mma_sfb, - tma_atom_a, tma_tensor_a, tma_atom_b, tma_tensor_b, - tma_atom_sfa, tma_tensor_sfa, tma_atom_sfb, tma_tensor_sfb, - tma_atom_c, tma_tensor_c, - self.cluster_layout_vmnk, self.cluster_layout_sfb_vmnk, - self.a_smem_layout_staged, self.b_smem_layout_staged, - self.sfa_smem_layout_staged, self.sfb_smem_layout_staged, - self.c_smem_layout_staged, - self.epi_tile, - tile_sched_params, - M, N, K, gsa, gsb, - ).launch( - grid=grid, block=[self.threads_per_cta, 1, 1], - cluster=(*self.cluster_shape_mn, 1), - stream=stream, min_blocks_per_mp=1, - ) - - cute.compile(_compiled_fn, mat_a, mat_b, scale_a, scale_b, mat_c) + # Launch kernel directly (same as fused_swiglu pattern) + self._kernel( + tiled_mma, tiled_mma_sfb, + tma_atom_a, tma_tensor_a, tma_atom_b, tma_tensor_b, + tma_atom_sfa, tma_tensor_sfa, tma_atom_sfb, tma_tensor_sfb, + tma_atom_c, tma_tensor_c, + self.cluster_layout_vmnk, self.cluster_layout_sfb_vmnk, + self.a_smem_layout_staged, self.b_smem_layout_staged, + self.sfa_smem_layout_staged, self.sfb_smem_layout_staged, + self.c_smem_layout_staged, + self.epi_tile, + tile_sched_params, + M, N, K, gsa, gsb, + ).launch( + grid=grid, block=[self.threads_per_cta, 1, 1], + cluster=(*self.cluster_shape_mn, 1), + stream=stream, min_blocks_per_mp=1, + ) # ----------------------------------------------------------------- # GPU kernel