From d8d241bf9f867ff270e7f26215fee7ee29023e9c Mon Sep 17 00:00:00 2001 From: Claude Date: Sat, 15 Aug 2026 12:34:01 +0000 Subject: [PATCH] =?UTF-8?q?fix:=20corex=5Fbatched=5Fgemm=5Fkernel=20?= =?UTF-8?q?=E2=80=94=20use=20OpClassTensorOp=20+=20arch::Cu10=20+=20FP32?= =?UTF-8?q?=20accumulator?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Root cause of 25ms (vs expected 2.5ms): 1. ElementAccumulator was half_t → now float (FP32 accumulation) 2. Missing OpClassTensorOp → was defaulting to OpClassSimt (CUDA cores only) 3. Missing arch::Cu10 → was defaulting to arch::Sm61 With these fixes it should use __ivcorex_matrix_mad_f32x4_f16x4 (TCU) same as moe_cutlass_batched.cu which benchmarked at 2.462ms. --- .../cuda/corex_batched_gemm_kernel.cu | 15 +++++++++++---- 1 file changed, 11 insertions(+), 4 deletions(-) diff --git a/ex_engine/xllm_kernels/cuda/corex_batched_gemm_kernel.cu b/ex_engine/xllm_kernels/cuda/corex_batched_gemm_kernel.cu index e8b273aa..ea8246ce 100644 --- a/ex_engine/xllm_kernels/cuda/corex_batched_gemm_kernel.cu +++ b/ex_engine/xllm_kernels/cuda/corex_batched_gemm_kernel.cu @@ -35,17 +35,24 @@ cudaError_t cutlass_batched_hgemm( using ElementA = cutlass::half_t; using ElementB = cutlass::half_t; using ElementC = cutlass::half_t; - using ElementAccumulator = cutlass::half_t; + using ElementAccumulator = float; using Gemm = cutlass::gemm::device::GemmBatched< ElementA, cutlass::layout::ColumnMajor, // A ElementB, cutlass::layout::ColumnMajor, // B ElementC, cutlass::layout::ColumnMajor, // C - ElementAccumulator // accumulator + ElementAccumulator, // accumulator = FP32 + cutlass::arch::OpClassTensorOp, // use TCU (not SIMT) + cutlass::arch::Cu10 // BI-V100 arch + // Defaults from DefaultGemmConfiguration: + // ThreadblockShape = <128, 128, 32> + // WarpShape = <32, 32, 32> + // InstructionShape = <16, 16, 16> + // Stages = 2 >; - ElementAccumulator alpha_val(1.0f); - ElementAccumulator beta_val(0.0f); + float alpha_val = 1.0f; + float beta_val = 0.0f; Gemm gemm_op;