diff --git a/dsv4/kernels/router/dense_router_decode.py b/dsv4/kernels/router/dense_router_decode.py index a7a1d74c..606a2f74 100644 --- a/dsv4/kernels/router/dense_router_decode.py +++ b/dsv4/kernels/router/dense_router_decode.py @@ -25,18 +25,16 @@ def dense_router_dispatch( """ N = hidden_states.shape[0] - # The CuTeDSL fused decode kernel has a TMA partition layout bug that - # causes cute.compile to fail after a long compilation attempt. - # TODO: fix the fused kernel (OperandMajorMode + local_tile coord mismatch) - # For now, the BF16 linear + activation_topk path is the production path. - # BF16 GEMM on Blackwell uses tensor cores via cuBLAS; the activation_topk - # kernel is a real CUDA kernel (not PyTorch reference). - # if N <= 64: - # try: - # _run_fused_decode(...) - # return - # except Exception: - # pass + if N <= 64: + try: + _run_fused_decode( + hidden_states, W_gate, e_bias, + routed_scaling_factor, top_k, + out_weights, out_ids, + ) + return + except Exception: + pass # fall through to prefill path _run_prefill_path( hidden_states, W_gate, e_bias, diff --git a/dsv4/kernels/router/dense_router_decode_kernel.py b/dsv4/kernels/router/dense_router_decode_kernel.py index e537406e..a175d5f3 100644 --- a/dsv4/kernels/router/dense_router_decode_kernel.py +++ b/dsv4/kernels/router/dense_router_decode_kernel.py @@ -101,18 +101,14 @@ class DenseRouterDecodeKernel: self.num_tmem_alloc_cols = utils.get_num_tmem_alloc_cols(tCtAcc_fake) def run(self, X, W_gate, e_bias, out_w, out_ids, M, E, K, scaling, top_k, stream=None): - # A operand (hidden_states) is (M, K) — MN-major layout - # B operand (W_gate) is (K, E) — K-major layout - self.a_major_mode = OperandMajorMode.MN - self.b_major_mode = OperandMajorMode.K - if stream is None: stream = cuda.CUstream(0) - # All MLIR-dependent setup (tiled_mma, TMA atoms, CuTe tensor conversion) - # must happen inside cute.compile context. This matches the MoE kernel pattern. @cute.jit def _compiled_fn(X, W_gate, e_bias, out_w, out_ids): + # Infer major modes from tensor layouts (same as MoE/grouped GEMM kernels) + self.a_major_mode = utils.LayoutEnum.from_tensor(X).mma_major_mode() + self.b_major_mode = utils.LayoutEnum.from_tensor(W_gate).mma_major_mode() self._setup_attributes() tiled_mma = self._tiled_mma atom_thr_size = cute.size(tiled_mma.thr_id.shape)