submission 371747
Joel🏴 · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 873 lines, June 9 Researcher Reciprocity License v1.0.
submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-modal-nvfp4-dual-gemm-371747?include=source"interfacepython
Compatibility
measured onNVIDIA B200
declared hardwareNVIDIA B200
architecturessm_100
dtypesfp8_e4m3, nvfp4
Benchmark evidence
1 measurement across 1 GPU, fastest first.
Operation / workload
Hardware
Latency
Rank
Observed
Reported · How evidence levels are derived →
Source and license
sourceavailable
revision digestsha256:261849e485462b912a2a5e9ad6d10bbaaf047173d11ebe5339a79728644bd94c
license declaredunknown
license concludedunknown
authorsJoel🏴
imported2026-08-26
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
fp4
const int K = A.size(1) * 2; // Convert from FP4 packed to element countfused-epilogue
WaitEpilogue,mbarrier
void mbarrier_init(int mbar_addr, int count) {shared-memory
extern __shared__ __align__(1024) char smem_ptr[];split-k
int SPLIT_K,stages = 5
LAUNCH( 4096, 128, 64, 256, 1, false, false, 5) // Keep baseline: BLOCK_N=64, NUM_STAGES=5tcgen05
asm volatile("tcgen05.cp.cta_group::1.32x128b.warpx4 [%0], %1;" :: "r"(taddr), "l"(s_desc));tile-n = 64
LAUNCH( 4096, 128, 64, 256, 1, false, false, 5) // Keep baseline: BLOCK_N=64, NUM_STAGES=5tma
asm volatile("cp.async.bulk.shared::cta.global.mbarrier::complete_tx::bytes.L2::cache_hint [%0], [%1], %2, [%3], %4;"vector-width = half2
reinterpret_cast<half2 *>(C_ptr + (row + 0) * N + col)[0] = __float22half2_rn({res[0], res[1]});Kernel source
submission.py873 lines
import torch
from torch.utils.cpp_extension import load_inline
# 1. CUDA Source: Utils
CUDA_SRC_UTILS = r"""
#ifndef UTILS_H
#define UTILS_H
#include <cuda_fp16.h>
#include <cuda_fp8.h> // Added for __nv_fp8_e4m3
#include <cuda.h>
#include <cstdint>
#include <cstdio>
#include <cstdlib>
#include <cuda_runtime.h> // Added for cudaSuccess, cudaGetErrorString etc.
// Define CHECK_CUDA and CHECK_CU macros here, outside the raw string
#define CHECK_CUDA(err) do { \
if (err != cudaSuccess) { \
fprintf(stderr, "CUDA Error: %s at %s:%d\n", cudaGetErrorString(err), __FILE__, __LINE__); \
exit(EXIT_FAILURE); \
} \
} while (0)
#define CHECK_CU(err) do { \
if (err != CUDA_SUCCESS) { \
const char *error_msg_ptr; \
cuGetErrorString(err, &error_msg_ptr); \
fprintf(stderr, "CU Error: %s at %s:%d\n", error_msg_ptr, __FILE__, __LINE__); \
exit(EXIT_FAILURE); \
} \
} while (0)
constexpr int WARP_SIZE = 32;
constexpr int MMA_K = 64; // 32 bytes
// https://github.com/NVIDIA/cutlass/blob/v4.3.2/include/cute/arch/copy_sm90_desc.hpp#L193-L197
constexpr uint64_t EVICT_NORMAL = 0x1000000000000000;
constexpr uint64_t EVICT_FIRST = 0x12F0000000000000;
constexpr uint64_t EVICT_LAST = 0x14F0000000000000;
enum ProfilerTag {
Setup = 0,
IssueTMA,
IssueMMA,
WaitTMA,
WaitMMA,
WaitMainloop,
WaitEpilogue,
Epilogue,
};
__device__ inline
int64_t globaltimer() {
int64_t t;
asm volatile("mov.u64 %0, %globaltimer;" : "=l"(t) :: "memory");
return t;
}
struct Profiler {
int64_t *data_ptr_;
int sm_id_;
int cnt_;
__device__
void init(int num_entries, int64_t *data_ptr, int bid) {
data_ptr_ = data_ptr + bid * (1 + num_entries * 4);
asm volatile("mov.u32 %0, %smid;\n" : "=r"(sm_id_));
cnt_ = 0;
}
__device__
void start(ProfilerTag tag) {
data_ptr_[1 + cnt_ * 4 + 0] = sm_id_;
data_ptr_[1 + cnt_ * 4 + 1] = tag;
data_ptr_[1 + cnt_ * 4 + 2] = globaltimer();
}
__device__
void stop() {
data_ptr_[1 + cnt_ * 4 + 3] = globaltimer() - data_ptr_[1 + cnt_ * 4 + 2];
cnt_ += 1;
}
__device__
void flush() {
data_ptr_[0] = cnt_;
}
};
__device__ inline
constexpr uint64_t desc_encode(uint64_t x) { return (x & 0x3'FFFFULL) >> 4ULL; };
// https://github.com/NVIDIA/cutlass/blob/v4.2.1/include/cute/arch/cluster_sm90.hpp#L180
__device__ inline
uint32_t elect_sync() {
#if defined(CUTE_ARCH_ELECT_ONE_SM90_ENABLED)
uint32_t pred = 0;
uint32_t laneid = 0;
asm volatile(
"{\n"
".reg .b32 %%rx;\n"
".reg .pred %%px;\n"
" elect.sync %%rx|%%px, %2;\n"
"@%%px mov.s32 %1, 1;\n"
" mov.s32 %0, %%rx;\n"
"}\n"
: "+r"(laneid), "+r"(pred)
: "r"(0xFFFFFFFF));
return pred;
#elif defined(__CUDA_ARCH__)
return (threadIdx.x % 32) == 0;
#else
return true;
#endif
}
__device__ inline
void mbarrier_init(int mbar_addr, int count) {
asm volatile("mbarrier.init.shared::cta.b64 [%0], %1;" :: "r"(mbar_addr), "r"(count));
}
// https://github.com/NVIDIA/cutlass/blob/v4.2.1/include/cutlass/arch/barrier.h#L408
__device__ inline
void mbarrier_wait(int mbar_addr, int phase) {
uint32_t ticks = 0x989680; // this is optional
asm volatile(
"{\n\t"
".reg .pred P1;\n\t"
"LAB_WAIT:\n\t"
"mbarrier.try_wait.parity.acquire.cta.shared::cta.b64 P1, [%0], %1, %2;\n\t"
"@P1 bra.uni DONE;\n\t"
"bra.uni LAB_WAIT;\n\t"
"DONE:\n\t"
"}"
:: "r"(mbar_addr), "r"(phase), "r"(ticks)
);
}
__device__ inline
void tma_gmem2smem(int dst, const void *src, int size, int mbar_addr, uint64_t cache_policy) {
asm volatile("cp.async.bulk.shared::cta.global.mbarrier::complete_tx::bytes.L2::cache_hint [%0], [%1], %2, [%3], %4;"
:: "r"(dst), "l"(src), "r"(size), "r"(mbar_addr), "l"(cache_policy));
}
__device__ inline void tma_3d_gmem2smem(int dst, const void *tmap_ptr, int x, int y, int z, int mbar_addr, uint64_t cache_policy) {
asm volatile("cp.async.bulk.tensor.3d.shared::cta.global.mbarrier::complete_tx::bytes.cta_group::1.L2::cache_hint "
"[%0], [%1, {%2, %3, %4}], [%5], %6;"
:: "r"(dst), "l"(tmap_ptr), "r"(x), "r"(y), "r"(z), "r"(mbar_addr), "l"(cache_policy)
: "memory");
}
__device__ inline void tcgen05_cp_nvfp4(int taddr, uint64_t s_desc) {
// .32x128b corresponds to (32, 16) 8-bit scale -> 1 MMA for nvfp4.
// .warpx4 duplicates data across 32-lane groups.
asm volatile("tcgen05.cp.cta_group::1.32x128b.warpx4 [%0], %1;" :: "r"(taddr), "l"(s_desc));
}
__device__ inline void tcgen05_mma_nvfp4(
uint64_t a_desc,
uint64_t b_desc,
uint32_t i_desc,
int scale_A_tmem,
int scale_B_tmem,
int enable_input_d
) {
const int d_tmem = 0; // assume
asm volatile(
"{\n\t"
".reg .pred p;\n\t" // predicate register enable-input-d
"setp.ne.b32 p, %6, 0;\n\t"
"tcgen05.mma.cta_group::1.kind::mxf4nvf4.block_scale.block16 [%0], %1, %2, %3, [%4], [%5], p;\n\t"
"}\n"
:: "r"(d_tmem), "l"(a_desc), "l"(b_desc), "r"(i_desc),
"r"(scale_A_tmem), "r"(scale_B_tmem), "r"(enable_input_d)
);
}
// see https://docs.nvidia.com/cuda/inline-ptx-assembly/index.html
struct SHAPE {
static constexpr char _32x32b[] = ".32x32b"; // 32x1 tile for each warp
static constexpr char _16x128b[] = ".16x128b"; // 16x4 tile
static constexpr char _16x256b[] = ".16x256b"; // 16x8 tile
};
struct NUM {
static constexpr char x4[] = ".x4";
static constexpr char x8[] = ".x8";
static constexpr char x16[] = ".x16";
static constexpr char x32[] = ".x32";
static constexpr char x64[] = ".x64";
static constexpr char x128[] = ".x128";
};
template <const char *SHAPE, const char *NUM>
__device__ inline
void tcgen05_ld_16regs(float *tmp, int row, int col) {
asm volatile("tcgen05.ld.sync.aligned%17%18.b32 "
"{ %0, %1, %2, %3, %4, %5, %6, %7, "
" %8, %9, %10, %11, %12, %13, %14, %15}, [%16];"
: "=f"(tmp[ 0]), "=f"(tmp[ 1]), "=f"(tmp[ 2]), "=f"(tmp[ 3]), "=f"(tmp[ 4]), "=f"(tmp[ 5]), "=f"(tmp[ 6]), "=f"(tmp[ 7]),
"=f"(tmp[ 8]), "=f"(tmp[ 9]), "=f"(tmp[10]), "=f"(tmp[11]), "=f"(tmp[12]), "=f"(tmp[13]), "=f"(tmp[14]), "=f"(tmp[15])
: "r"((row << 16) | col), "C"(SHAPE), "C"(NUM));
}
template <const char *SHAPE, const char *NUM>
__device__ inline
void tcgen05_ld_32regs(float *tmp, int row, int col) {
asm volatile("tcgen05.ld.sync.aligned%33%34.b32 "
"{ %0, %1, %2, %3, %4, %5, %6, %7, "
" %8, %9, %10, %11, %12, %13, %14, %15, "
" %16, %17, %18, %19, %20, %21, %22, %23, "
" %24, %25, %26, %27, %28, %29, %30, %31}, [%32];"
: "=f"(tmp[ 0]), "=f"(tmp[ 1]), "=f"(tmp[ 2]), "=f"(tmp[ 3]), "=f"(tmp[ 4]), "=f"(tmp[ 5]), "=f"(tmp[ 6]), "=f"(tmp[ 7]),
"=f"(tmp[ 8]), "=f"(tmp[ 9]), "=f"(tmp[10]), "=f"(tmp[11]), "=f"(tmp[12]), "=f"(tmp[13]), "=f"(tmp[14]), "=f"(tmp[15]),
"=f"(tmp[16]), "=f"(tmp[17]), "=f"(tmp[18]), "=f"(tmp[19]), "=f"(tmp[20]), "=f"(tmp[21]), "=f"(tmp[22]), "=f"(tmp[23]),
"=f"(tmp[24]), "=f"(tmp[25]), "=f"(tmp[26]), "=f"(tmp[27]), "=f"(tmp[28]), "=f"(tmp[29]), "=f"(tmp[30]), "=f"(tmp[31])
: "r"((row << 16) | col), "C"(SHAPE), "C"(NUM));
}
template <const char *SHAPE, const char *NUM>
__device__ inline
void tcgen05_ld_64regs(float *tmp, int row, int col) {
asm volatile("tcgen05.ld.sync.aligned%65%66.b32 "
"{ %0, %1, %2, %3, %4, %5, %6, %7, "
" %8, %9, %10, %11, %12, %13, %14, %15, "
" %16, %17, %18, %19, %20, %21, %22, %23, "
" %24, %25, %26, %27, %28, %29, %30, %31, "
" %32, %33, %34, %35, %36, %37, %38, %39, "
" %40, %41, %42, %43, %44, %45, %46, %47, "
" %48, %49, %50, %51, %52, %53, %54, %55, "
" %56, %57, %58, %59, %60, %61, %62, %63}, [%64];"
: "=f"(tmp[ 0]), "=f"(tmp[ 1]), "=f"(tmp[ 2]), "=f"(tmp[ 3]), "=f"(tmp[ 4]), "=f"(tmp[ 5]), "=f"(tmp[ 6]), "=f"(tmp[ 7]),
"=f"(tmp[ 8]), "=f"(tmp[ 9]), "=f"(tmp[10]), "=f"(tmp[11]), "=f"(tmp[12]), "=f"(tmp[13]), "=f"(tmp[14]), "=f"(tmp[15]),
"=f"(tmp[16]), "=f"(tmp[17]), "=f"(tmp[18]), "=f"(tmp[19]), "=f"(tmp[20]), "=f"(tmp[21]), "=f"(tmp[22]), "=f"(tmp[23]),
"=f"(tmp[24]), "=f"(tmp[25]), "=f"(tmp[26]), "=f"(tmp[27]), "=f"(tmp[28]), "=f"(tmp[29]), "=f"(tmp[30]), "=f"(tmp[31]),
"=f"(tmp[32]), "=f"(tmp[33]), "=f"(tmp[34]), "=f"(tmp[35]), "=f"(tmp[36]), "=f"(tmp[37]), "=f"(tmp[38]), "=f"(tmp[39]),
"=f"(tmp[40]), "=f"(tmp[41]), "=f"(tmp[42]), "=f"(tmp[43]), "=f"(tmp[44]), "=f"(tmp[45]), "=f"(tmp[46]), "=f"(tmp[47]),
"=f"(tmp[48]), "=f"(tmp[49]), "=f"(tmp[50]), "=f"(tmp[51]), "=f"(tmp[52]), "=f"(tmp[53]), "=f"(tmp[54]), "=f"(tmp[55]),
"=f"(tmp[56]), "=f"(tmp[57]), "=f"(tmp[58]), "=f"(tmp[59]), "=f"(tmp[60]), "=f"(tmp[61]), "=f"(tmp[62]), "=f"(tmp[63])
: "r"((row << 16) | col), "C"(SHAPE), "C"(NUM));
}
template <const char *SHAPE, const char *NUM>
__device__ inline
void tcgen05_ld_128regs(float *tmp, int row, int col) {
asm volatile("tcgen05.ld.sync.aligned%129%130.b32 "
"{ %0, %1, %2, %3, %4, %5, %6, %7, "
" %8, %9, %10, %11, %12, %13, %14, %15, "
" %16, %17, %18, %19, %20, %21, %22, %23, "
" %24, %25, %26, %27, %28, %29, %30, %31, "
" %32, %33, %34, %35, %36, %37, %38, %39, "
" %40, %41, %42, %43, %44, %45, %46, %47, "
" %48, %49, %50, %51, %52, %53, %54, %55, "
" %56, %57, %58, %59, %60, %61, %62, %63, "
" %64, %65, %66, %67, %68, %69, %70, %71, "
" %72, %73, %74, %75, %76, %77, %78, %79, "
" %80, %81, %82, %83, %84, %85, %86, %87, "
" %88, %89, %90, %91, %92, %93, %94, %95, "
" %96, %97, %98, %99,%100,%101,%102,%103, "
"%104,%105,%106,%107,%108,%109,%110,%111, "
"%112,%113,%114,%115,%116,%117,%118,%119, "
"%120,%121,%122,%123,%124,%125,%126,%127}, [%128];"
: "=f"(tmp[ 0]), "=f"(tmp[ 1]), "=f"(tmp[ 2]), "=f"(tmp[ 3]), "=f"(tmp[ 4]), "=f"(tmp[ 5]), "=f"(tmp[ 6]), "=f"(tmp[ 7]),
"=f"(tmp[ 8]), "=f"(tmp[ 9]), "=f"(tmp[10]), "=f"(tmp[11]), "=f"(tmp[12]), "=f"(tmp[13]), "=f"(tmp[14]), "=f"(tmp[15]),
"=f"(tmp[16]), "=f"(tmp[17]), "=f"(tmp[18]), "=f"(tmp[19]), "=f"(tmp[20]), "=f"(tmp[21]), "=f"(tmp[22]), "=f"(tmp[23]),
"=f"(tmp[24]), "=f"(tmp[25]), "=f"(tmp[26]), "=f"(tmp[27]), "=f"(tmp[28]), "=f"(tmp[29]), "=f"(tmp[30]), "=f"(tmp[31]),
"=f"(tmp[32]), "=f"(tmp[33]), "=f"(tmp[34]), "=f"(tmp[35]), "=f"(tmp[36]), "=f"(tmp[37]), "=f"(tmp[38]), "=f"(tmp[39]),
"=f"(tmp[40]), "=f"(tmp[41]), "=f"(tmp[42]), "=f"(tmp[43]), "=f"(tmp[44]), "=f"(tmp[45]), "=f"(tmp[46]), "=f"(tmp[47]),
"=f"(tmp[48]), "=f"(tmp[49]), "=f"(tmp[50]), "=f"(tmp[51]), "=f"(tmp[52]), "=f"(tmp[53]), "=f"(tmp[54]), "=f"(tmp[55]),
"=f"(tmp[56]), "=f"(tmp[57]), "=f"(tmp[58]), "=f"(tmp[59]), "=f"(tmp[60]), "=f"(tmp[61]), "=f"(tmp[62]), "=f"(tmp[63]),
"=f"(tmp[64]), "=f"(tmp[65]), "=f"(tmp[66]), "=f"(tmp[67]), "=f"(tmp[68]), "=f"(tmp[69]), "=f"(tmp[70]), "=f"(tmp[71]),
"=f"(tmp[72]), "=f"(tmp[73]), "=f"(tmp[74]), "=f"(tmp[75]), "=f"(tmp[76]), "=f"(tmp[77]), "=f"(tmp[78]), "=f"(tmp[79]),
"=f"(tmp[80]), "=f"(tmp[81]), "=f"(tmp[82]), "=f"(tmp[83]), "=f"(tmp[84]), "=f"(tmp[85]), "=f"(tmp[86]), "=f"(tmp[87]),
"=f"(tmp[88]), "=f"(tmp[89]), "=f"(tmp[90]), "=f"(tmp[91]), "=f"(tmp[92]), "=f"(tmp[93]), "=f"(tmp[94]), "=f"(tmp[95]),
"=f"(tmp[96]), "=f"(tmp[97]), "=f"(tmp[98]), "=f"(tmp[99]), "=f"(tmp[100]),"=f"(tmp[101]),"=f"(tmp[102]),"=f"(tmp[103]),
"=f"(tmp[104]),"=f"(tmp[105]),"=f"(tmp[106]),"=f"(tmp[107]),"=f"(tmp[108]),"=f"(tmp[109]),"=f"(tmp[110]),"=f"(tmp[111]),
"=f"(tmp[112]),"=f"(tmp[113]),"=f"(tmp[114]),"=f"(tmp[115]),"=f"(tmp[116]),"=f"(tmp[117]),"=f"(tmp[118]),"=f"(tmp[119]),
"=f"(tmp[120]),"=f"(tmp[121]),"=f"(tmp[122]),"=f"(tmp[123]),"=f"(tmp[124]),"=f"(tmp[125]),"=f"(tmp[126]),"=f"(tmp[127])
: "r"((row << 16) | col), "C"(SHAPE), "C"(NUM));
}
__device__ inline void tcgen05_ld_32x32bx32(float *tmp, int row, int col) { tcgen05_ld_32regs<SHAPE::_32x32b, NUM::x32>(tmp, row, col); }
__device__ inline void tcgen05_ld_32x32bx64(float *tmp, int row, int col) { tcgen05_ld_64regs<SHAPE::_32x32b, NUM::x64>(tmp, row, col); }
__device__ inline void tcgen05_ld_32x32bx128(float *tmp, int row, int col) { tcgen05_ld_128regs<SHAPE::_32x32b, NUM::x128>(tmp, row, col); }
__device__ inline void tcgen05_ld_16x128bx8(float *tmp, int row, int col) { tcgen05_ld_16regs<SHAPE::_16x128b, NUM::x8>(tmp, row, col); }
__device__ inline void tcgen05_ld_16x128bx16(float *tmp, int row, int col) { tcgen05_ld_32regs<SHAPE::_16x128b, NUM::x16>(tmp, row, col); }
__device__ inline void tcgen05_ld_16x128bx32(float *tmp, int row, int col) { tcgen05_ld_64regs<SHAPE::_16x128b, NUM::x32>(tmp, row, col); }
__device__ inline void tcgen05_ld_16x256bx4(float *tmp, int row, int col) { tcgen05_ld_16regs<SHAPE::_16x256b, NUM::x4>(tmp, row, col); }
__device__ inline void tcgen05_ld_16x256bx8(float *tmp, int row, int col) { tcgen05_ld_32regs<SHAPE::_16x256b, NUM::x8>(tmp, row, col); }
__device__ inline void tcgen05_ld_16x256bx16(float *tmp, int row, int col) { tcgen05_ld_64regs<SHAPE::_16x256b, NUM::x16>(tmp, row, col); }
inline void init_AB_tmap(
CUtensorMap *tmap,
const char *ptr,
uint64_t global_height, uint64_t global_width,
uint32_t shared_height, uint32_t shared_width
) {
constexpr uint32_t rank = 3;
uint64_t globalDim[rank] = {256, global_height, global_width / 256};
uint64_t globalStrides[rank-1] = {global_width / 2, 128}; // in bytes
uint32_t boxDim[rank] = {256, shared_height, shared_width / 256};
uint32_t elementStrides[rank] = {1, 1, 1};
auto err = cuTensorMapEncodeTiled(
tmap,
CUtensorMapDataType::CU_TENSOR_MAP_DATA_TYPE_16U4_ALIGN8B,
rank,
(void *)ptr,
globalDim,
globalStrides,
boxDim,
elementStrides,
CUtensorMapInterleave::CU_TENSOR_MAP_INTERLEAVE_NONE,
CUtensorMapSwizzle::CU_TENSOR_MAP_SWIZZLE_128B,
CUtensorMapL2promotion::CU_TENSOR_MAP_L2_PROMOTION_NONE,
CUtensorMapFloatOOBfill::CU_TENSOR_MAP_FLOAT_OOB_FILL_NONE
);
CHECK_CU(err);
}
#endif // UTILS_H
"""
# 2. CUDA Source: Kernel (minus includes)
CUDA_SRC_KERNEL = r"""
#include <algorithm>
__device__ inline float silu_exact(float x) {
// Exact SiLU: x / (1 + exp(-x)) = x * sigmoid(x)
return x / (1.0f + expf(-x));
}
template <
int K,
int BLOCK_M,
int BLOCK_N,
int BLOCK_K,
int SPLIT_K,
bool C_N_MAJOR,
int NUM_STAGES
>
__global__
__launch_bounds__(BLOCK_M + 2 * WARP_SIZE)
void kernel_v4_dual(
const __grid_constant__ CUtensorMap A_tmap,
const __grid_constant__ CUtensorMap B1_tmap,
const __grid_constant__ CUtensorMap B2_tmap,
const char *SFA_ptr,
const char *SFB1_ptr,
const char *SFB2_ptr,
half *C_ptr,
float *buf_ptr,
int M, int N
) {
const int tid = threadIdx.x;
const int bid_k = blockIdx.x;
const int bid = blockIdx.y;
const int lane_id = tid % WARP_SIZE;
const int warp_id = tid / WARP_SIZE;
const int grid_m = M / BLOCK_M;
const int grid_n = N / BLOCK_N;
const int bid_m = bid / grid_n;
const int bid_n = bid % grid_n;
const int off_m = bid_m * BLOCK_M;
const int off_n = bid_n * BLOCK_N;
constexpr int NUM_WARPS = BLOCK_M / WARP_SIZE + 2;
// set up smem
extern __shared__ __align__(1024) char smem_ptr[];
const int smem = static_cast<int>(__cvta_generic_to_shared(smem_ptr));
constexpr int A_size = BLOCK_M * BLOCK_K / 2;
constexpr int B_size = BLOCK_N * BLOCK_K / 2;
constexpr int SFA_size = 128 * BLOCK_K / 16; // always copy 128xBLOCK_K/16
constexpr int SFB_size = 128 * BLOCK_K / 16;
constexpr int STAGE_SIZE = A_size + B_size + SFA_size + SFB_size;
// Intermediate accumulator storage (Phase 1 result)
float *acc1_smem = reinterpret_cast<float*>(smem_ptr + STAGE_SIZE * NUM_STAGES);
// set up mbarriers and tmem
#pragma nv_diag_suppress static_var_with_dynamic_init
__shared__ int64_t mbars[NUM_STAGES * 2 + 1];
const int tma_mbar_addr = static_cast<int>(__cvta_generic_to_shared(mbars));
const int mma_mbar_addr = tma_mbar_addr + NUM_STAGES * 8;
const int mainloop_mbar_addr = mma_mbar_addr + NUM_STAGES * 8;
constexpr int SFA_tmem = BLOCK_N;
constexpr int SFB_tmem = SFA_tmem + 4 * (BLOCK_K / MMA_K);
constexpr int num_iters = K / BLOCK_K / SPLIT_K;
// ========================================================================
// PHASE 1: GEMM (A @ B1)
// ========================================================================
if (warp_id == 0 && elect_sync()) {
for (int i = 0; i < NUM_STAGES * 2 + 1; i++)
mbarrier_init(tma_mbar_addr + i * 8, 1);
asm volatile("fence.mbarrier_init.release.cluster;");
}
else if (warp_id == 1) {
asm volatile("tcgen05.alloc.cta_group::1.sync.aligned.shared::cta.b32 [%0], %1;" :: "r"(smem), "r"(BLOCK_N * 2));
}
__syncthreads();
if (warp_id == NUM_WARPS - 2 && elect_sync()) {
// TMA warp
uint64_t cache_A, cache_B;
if (M > N) {
cache_A = EVICT_FIRST;
cache_B = EVICT_LAST;
} else {
cache_A = EVICT_LAST;
cache_B = EVICT_FIRST;
}
auto issue_tma = [&](int iter_k, int stage_id) {
const int mbar_addr = tma_mbar_addr + stage_id * 8;
const int A_smem = smem + stage_id * STAGE_SIZE;
const int B_smem = A_smem + A_size;
const int SFA_smem = B_smem + B_size;
const int SFB_smem = SFA_smem + SFA_size;
const int off_k = SPLIT_K == 1 ? iter_k * BLOCK_K : (iter_k * SPLIT_K + bid_k) * BLOCK_K;
tma_3d_gmem2smem(A_smem, &A_tmap, 0, off_m, off_k / 256, mbar_addr, cache_A);
tma_3d_gmem2smem(B_smem, &B1_tmap, 0, off_n, off_k / 256, mbar_addr, cache_B); // Load B1
const int rest_k = K / 16 / 4;
const char *SFA_src = SFA_ptr + ((off_m / 128) * rest_k + off_k / (16 * 4)) * 512;
const char *SFB_src = SFB1_ptr + ((off_n / 128) * rest_k + off_k / (16 * 4)) * 512; // Load SFB1
tma_gmem2smem(SFA_smem, SFA_src, SFA_size, mbar_addr, cache_A);
tma_gmem2smem(SFB_smem, SFB_src, SFB_size, mbar_addr, cache_B);
asm volatile("mbarrier.arrive.expect_tx.release.cta.shared::cta.b64 _, [%0], %1;"
:: "r"(mbar_addr), "r"(STAGE_SIZE) : "memory");
};
for (int iter_k = 0; iter_k < NUM_STAGES; iter_k++)
issue_tma(iter_k, iter_k);
for (int iter_k = NUM_STAGES; iter_k < num_iters; iter_k++) {
const int stage_id = iter_k % NUM_STAGES;
const int mma_phase = (iter_k / NUM_STAGES - 1) % 2;
mbarrier_wait(mma_mbar_addr + stage_id * 8, mma_phase);
issue_tma(iter_k, stage_id);
}
}
else if (warp_id == NUM_WARPS - 1 && elect_sync()) {
// MMA warp
constexpr int MMA_N = BLOCK_N;
constexpr int MMA_M = 128;
constexpr uint32_t i_desc = (1U << 7U)
| (1U << 10U)
| ((uint32_t)MMA_N >> 3U << 17U)
| ((uint32_t)MMA_M >> 7U << 27U);
for (int iter_k = 0; iter_k < num_iters; iter_k++) {
const int stage_id = iter_k % NUM_STAGES;
const int tma_phase = (iter_k / NUM_STAGES) % 2;
mbarrier_wait(tma_mbar_addr + stage_id * 8, tma_phase);
const int A_smem = smem + stage_id * STAGE_SIZE;
const int B_smem = A_smem + A_size;
const int SFA_smem = B_smem + B_size;
const int SFB_smem = SFA_smem + SFA_size;
auto make_desc_AB = [](int addr) -> uint64_t {
const int SBO = 8 * 128;
return desc_encode(addr) | (desc_encode(SBO) << 32ULL) | (1ULL << 46ULL) | (2ULL << 61ULL);
};
auto make_desc_SF = [](int addr) -> uint64_t {
const int SBO = 8 * 16;
return desc_encode(addr) | (desc_encode(SBO) << 32ULL) | (1ULL << 46ULL);
};
constexpr uint64_t SF_desc = make_desc_SF(0);
const uint64_t SFA_desc = SF_desc + ((uint64_t)SFA_smem >> 4ULL);
const uint64_t SFB_desc = SF_desc + ((uint64_t)SFB_smem >> 4ULL);
for (int k = 0; k < BLOCK_K / MMA_K; k++) {
uint64_t sfa_desc = SFA_desc + (uint64_t)k * (512ULL >> 4ULL);
uint64_t sfb_desc = SFB_desc + (uint64_t)k * (512ULL >> 4ULL);
tcgen05_cp_nvfp4(SFA_tmem + k * 4, sfa_desc);
tcgen05_cp_nvfp4(SFB_tmem + k * 4, sfb_desc);
}
for (int k1 = 0; k1 < BLOCK_K / 256; k1++)
for (int k2 = 0; k2 < 256 / MMA_K; k2++) {
uint64_t a_desc = make_desc_AB(A_smem + k1 * BLOCK_M * 128 + k2 * 32);
uint64_t b_desc = make_desc_AB(B_smem + k1 * BLOCK_N * 128 + k2 * 32);
int k_sf = k1 * 4 + k2;
const int scale_A_tmem = SFA_tmem + k_sf * 4 + (bid_m % (128 / BLOCK_M)) * (BLOCK_M / 32);
const int scale_B_tmem = SFB_tmem + k_sf * 4 + (bid_n % (128 / BLOCK_N)) * (BLOCK_N / 32);
const int enable_input_d = (k1 == 0 && k2 == 0) ? iter_k : 1;
tcgen05_mma_nvfp4(a_desc, b_desc, i_desc, scale_A_tmem, scale_B_tmem, enable_input_d);
}
asm volatile("tcgen05.commit.cta_group::1.mbarrier::arrive::one.shared::cluster.b64 [%0];"
:: "r"(mma_mbar_addr + stage_id * 8) : "memory");
}
asm volatile("tcgen05.commit.cta_group::1.mbarrier::arrive::one.shared::cluster.b64 [%0];"
:: "r"(mainloop_mbar_addr) : "memory");
}
else if (tid < BLOCK_M) {
// Save Acc1 to SMEM
mbarrier_wait(mainloop_mbar_addr, 0); // Phase 0
asm volatile("tcgen05.fence::after_thread_sync;");
constexpr int WIDTH = std::min(BLOCK_N, 64);
for (int n = 0; n < BLOCK_N / WIDTH; n++) {
float tmp[WIDTH];
if constexpr (WIDTH == 128) tcgen05_ld_32x32bx128(tmp, warp_id * 32, n * WIDTH);
else if constexpr (WIDTH == 64) tcgen05_ld_32x32bx64(tmp, warp_id * 32, n * WIDTH);
else if constexpr (WIDTH == 32) tcgen05_ld_32x32bx32(tmp, warp_id * 32, n * WIDTH);
asm volatile("tcgen05.wait::ld.sync.aligned;");
for (int i = 0; i < WIDTH; i++) {
int row = tid; // tid < BLOCK_M (128)
int col = n * WIDTH + i;
acc1_smem[row * BLOCK_N + col] = tmp[i];
}
}
}
__syncthreads(); // Sync before Phase 2
// ========================================================================
// PHASE 2: GEMM (A @ B2)
// ========================================================================
// Reset barriers for Phase 2
if (warp_id == 0 && elect_sync()) {
for (int i = 0; i < NUM_STAGES * 2 + 1; i++)
mbarrier_init(tma_mbar_addr + i * 8, 1);
asm volatile("fence.mbarrier_init.release.cluster;");
}
__syncthreads();
if (warp_id == NUM_WARPS - 2 && elect_sync()) {
// TMA warp - Load B2
uint64_t cache_A, cache_B;
if (M > N) {
cache_A = EVICT_FIRST;
cache_B = EVICT_LAST;
} else {
cache_A = EVICT_LAST;
cache_B = EVICT_FIRST;
}
auto issue_tma = [&](int iter_k, int stage_id) {
const int mbar_addr = tma_mbar_addr + stage_id * 8;
const int A_smem = smem + stage_id * STAGE_SIZE;
const int B_smem = A_smem + A_size;
const int SFA_smem = B_smem + B_size;
const int SFB_smem = SFA_smem + SFA_size;
const int off_k = SPLIT_K == 1 ? iter_k * BLOCK_K : (iter_k * SPLIT_K + bid_k) * BLOCK_K;
tma_3d_gmem2smem(A_smem, &A_tmap, 0, off_m, off_k / 256, mbar_addr, cache_A);
tma_3d_gmem2smem(B_smem, &B2_tmap, 0, off_n, off_k / 256, mbar_addr, cache_B); // Load B2
const int rest_k = K / 16 / 4;
const char *SFA_src = SFA_ptr + ((off_m / 128) * rest_k + off_k / (16 * 4)) * 512;
const char *SFB_src = SFB2_ptr + ((off_n / 128) * rest_k + off_k / (16 * 4)) * 512; // Load SFB2
tma_gmem2smem(SFA_smem, SFA_src, SFA_size, mbar_addr, cache_A);
tma_gmem2smem(SFB_smem, SFB_src, SFB_size, mbar_addr, cache_B);
asm volatile("mbarrier.arrive.expect_tx.release.cta.shared::cta.b64 _, [%0], %1;"
:: "r"(mbar_addr), "r"(STAGE_SIZE) : "memory");
};
for (int iter_k = 0; iter_k < NUM_STAGES; iter_k++)
issue_tma(iter_k, iter_k);
for (int iter_k = NUM_STAGES; iter_k < num_iters; iter_k++) {
const int stage_id = iter_k % NUM_STAGES;
const int mma_phase = (iter_k / NUM_STAGES - 1) % 2;
mbarrier_wait(mma_mbar_addr + stage_id * 8, mma_phase);
issue_tma(iter_k, stage_id);
}
}
else if (warp_id == NUM_WARPS - 1 && elect_sync()) {
// MMA warp
constexpr int MMA_N = BLOCK_N;
constexpr int MMA_M = 128;
constexpr uint32_t i_desc = (1U << 7U)
| (1U << 10U)
| ((uint32_t)MMA_N >> 3U << 17U)
| ((uint32_t)MMA_M >> 7U << 27U);
for (int iter_k = 0; iter_k < num_iters; iter_k++) {
const int stage_id = iter_k % NUM_STAGES;
const int tma_phase = (iter_k / NUM_STAGES) % 2;
mbarrier_wait(tma_mbar_addr + stage_id * 8, tma_phase);
const int A_smem = smem + stage_id * STAGE_SIZE;
const int B_smem = A_smem + A_size;
const int SFA_smem = B_smem + B_size;
const int SFB_smem = SFA_smem + SFA_size;
auto make_desc_AB = [](int addr) -> uint64_t {
const int SBO = 8 * 128;
return desc_encode(addr) | (desc_encode(SBO) << 32ULL) | (1ULL << 46ULL) | (2ULL << 61ULL);
};
auto make_desc_SF = [](int addr) -> uint64_t {
const int SBO = 8 * 16;
return desc_encode(addr) | (desc_encode(SBO) << 32ULL) | (1ULL << 46ULL);
};
constexpr uint64_t SF_desc = make_desc_SF(0);
const uint64_t SFA_desc = SF_desc + ((uint64_t)SFA_smem >> 4ULL);
const uint64_t SFB_desc = SF_desc + ((uint64_t)SFB_smem >> 4ULL);
for (int k = 0; k < BLOCK_K / MMA_K; k++) {
uint64_t sfa_desc = SFA_desc + (uint64_t)k * (512ULL >> 4ULL);
uint64_t sfb_desc = SFB_desc + (uint64_t)k * (512ULL >> 4ULL);
tcgen05_cp_nvfp4(SFA_tmem + k * 4, sfa_desc);
tcgen05_cp_nvfp4(SFB_tmem + k * 4, sfb_desc);
}
for (int k1 = 0; k1 < BLOCK_K / 256; k1++)
for (int k2 = 0; k2 < 256 / MMA_K; k2++) {
uint64_t a_desc = make_desc_AB(A_smem + k1 * BLOCK_M * 128 + k2 * 32);
uint64_t b_desc = make_desc_AB(B_smem + k1 * BLOCK_N * 128 + k2 * 32);
int k_sf = k1 * 4 + k2;
const int scale_A_tmem = SFA_tmem + k_sf * 4 + (bid_m % (128 / BLOCK_M)) * (BLOCK_M / 32);
const int scale_B_tmem = SFB_tmem + k_sf * 4 + (bid_n % (128 / BLOCK_N)) * (BLOCK_N / 32);
const int enable_input_d = (k1 == 0 && k2 == 0) ? iter_k : 1;
tcgen05_mma_nvfp4(a_desc, b_desc, i_desc, scale_A_tmem, scale_B_tmem, enable_input_d);
}
asm volatile("tcgen05.commit.cta_group::1.mbarrier::arrive::one.shared::cluster.b64 [%0];"
:: "r"(mma_mbar_addr + stage_id * 8) : "memory");
}
asm volatile("tcgen05.commit.cta_group::1.mbarrier::arrive::one.shared::cluster.b64 [%0];"
:: "r"(mainloop_mbar_addr) : "memory");
}
else if (tid < BLOCK_M) {
// Epilogue - Compute SiLU(Acc1) * Acc2 and store
mbarrier_wait(mainloop_mbar_addr, 0);
asm volatile("tcgen05.fence::after_thread_sync;");
if constexpr (C_N_MAJOR) {
for (int m = 0; m < 32 / 16; m++) {
float tmp_acc2[BLOCK_N / 2];
if constexpr (BLOCK_N == 128) tcgen05_ld_16x256bx16(tmp_acc2, warp_id * 32 + m * 16, 0);
else if constexpr (BLOCK_N == 64) tcgen05_ld_16x256bx8(tmp_acc2, warp_id * 32 + m * 16, 0);
else if constexpr (BLOCK_N == 32) tcgen05_ld_16x256bx4(tmp_acc2, warp_id * 32 + m * 16, 0);
asm volatile("tcgen05.wait::ld.sync.aligned;");
for (int i = 0; i < BLOCK_N / 8; i++) {
const int row = off_m + warp_id * 32 + m * 16 + lane_id / 4;
const int col = off_n + i * 8 + (lane_id % 4) * 2;
int local_row = warp_id * 32 + m * 16 + lane_id / 4;
int local_col = i * 8 + (lane_id % 4) * 2;
float v_acc2[4] = {tmp_acc2[i*4+0], tmp_acc2[i*4+1], tmp_acc2[i*4+2], tmp_acc2[i*4+3]};
float v_acc1[4];
v_acc1[0] = acc1_smem[(local_row + 0) * BLOCK_N + local_col + 0];
v_acc1[1] = acc1_smem[(local_row + 0) * BLOCK_N + local_col + 1];
v_acc1[2] = acc1_smem[(local_row + 8) * BLOCK_N + local_col + 0];
v_acc1[3] = acc1_smem[(local_row + 8) * BLOCK_N + local_col + 1];
float res[4];
for(int k=0; k<4; ++k) res[k] = silu_exact(v_acc1[k]) * v_acc2[k];
if constexpr (SPLIT_K == 1) {
reinterpret_cast<half2 *>(C_ptr + (row + 0) * N + col)[0] = __float22half2_rn({res[0], res[1]});
reinterpret_cast<half2 *>(C_ptr + (row + 8) * N + col)[0] = __float22half2_rn({res[2], res[3]});
} else {
atomicAdd(reinterpret_cast<float2 *>(buf_ptr + (row + 0) * N + col), float2({res[0], res[1]}));
atomicAdd(reinterpret_cast<float2 *>(buf_ptr + (row + 8) * N + col), float2({res[2], res[3]}));
}
}
}
} else {
constexpr int WIDTH = std::min(BLOCK_N, 64);
for (int n = 0; n < BLOCK_N / WIDTH; n++) {
float tmp_acc2[WIDTH];
if constexpr (WIDTH == 128) tcgen05_ld_32x32bx128(tmp_acc2, warp_id * 32, n * WIDTH);
else if constexpr (WIDTH == 64) tcgen05_ld_32x32bx64(tmp_acc2, warp_id * 32, n * WIDTH);
else if constexpr (WIDTH == 32) tcgen05_ld_32x32bx32(tmp_acc2, warp_id * 32, n * WIDTH);
asm volatile("tcgen05.wait::ld.sync.aligned;");
for (int i = 0; i < WIDTH; i++) {
const int local_m = tid;
const int local_n = n * WIDTH + i;
const int row = off_m + local_m;
const int col = off_n + local_n;
float val_acc1 = acc1_smem[local_m * BLOCK_N + local_n];
float val_acc2 = tmp_acc2[i];
float result = silu_exact(val_acc1) * val_acc2;
if constexpr (SPLIT_K == 1)
C_ptr[row * N + col] = __float2half(result);
else
atomicAdd(buf_ptr + row * N + col, result);
}
}
}
asm volatile("bar.sync 1, %0;" :: "r"(BLOCK_M) : "memory");
if (warp_id == 0)
asm volatile("tcgen05.dealloc.cta_group::1.sync.aligned.b32 %0, %1;" :: "r"(0), "r"(BLOCK_N * 2));
}
}
template <
int K,
int BLOCK_M,
int BLOCK_N,
int BLOCK_K,
int SPLIT_K,
bool SWAP_AB,
bool C_N_MAJOR,
int NUM_STAGES
>
void dual_gemm_launch(
int M, int N,
const char* A_ptr,
const char* B1_ptr,
const char* B2_ptr,
const char* SFA_ptr,
const char* SFB1_ptr,
const char* SFB2_ptr,
half* C_ptr,
float* buf_ptr
) {
static_assert(BLOCK_K % 256 == 0);
int new_M = M;
int new_N = N;
// No SWAP_AB for dual GEMM - keep original matrix layout
CUtensorMap A_tmap, B1_tmap, B2_tmap;
init_AB_tmap(&A_tmap, A_ptr, new_M, K, BLOCK_M, BLOCK_K);
init_AB_tmap(&B1_tmap, B1_ptr, new_N, K, BLOCK_N, BLOCK_K);
init_AB_tmap(&B2_tmap, B2_ptr, new_N, K, BLOCK_N, BLOCK_K);
dim3 grid(SPLIT_K, (new_M / BLOCK_M) * (new_N / BLOCK_N));
int tb_size = BLOCK_M + 2 * WARP_SIZE;
int AB_size = (BLOCK_M + BLOCK_N) * (BLOCK_K / 2);
int SFAB_size = 128 * (BLOCK_K / 16) * 2;
int smem_size = (AB_size + SFAB_size) * NUM_STAGES + BLOCK_M * BLOCK_N * 4; // Add Acc1 space
auto this_kernel = kernel_v4_dual<K, BLOCK_M, BLOCK_N, BLOCK_K, SPLIT_K, C_N_MAJOR != SWAP_AB, NUM_STAGES>;
if (smem_size > 48000)
CHECK_CUDA(cudaFuncSetAttribute(this_kernel, cudaFuncAttributeMaxDynamicSharedMemorySize, smem_size));
this_kernel<<<grid, tb_size, smem_size>>>(A_tmap, B1_tmap, B2_tmap, SFA_ptr, SFB1_ptr, SFB2_ptr, C_ptr, buf_ptr, new_M, new_N);
CHECK_CUDA(cudaGetLastError());
}
""";
# 3. CUDA Source: Wrapper
CUDA_SRC_WRAPPER = r"""
#include <cuda.h>
#include <cuda_runtime.h>
#include <cuda_fp16.h>
#include <stdio.h>
#include <torch/extension.h>
at::Tensor dual_gemm_silu(
const at::Tensor& A,
const at::Tensor& B1,
const at::Tensor& B2,
const at::Tensor& SFA,
const at::Tensor& SFB1,
const at::Tensor& SFB2,
at::Tensor& C,
at::Tensor& buf
) {
const int M = A.size(0);
const int K = A.size(1) * 2; // Convert from FP4 packed to element count
const int N = B1.size(0);
#define LAUNCH(K_, BLOCK_M, BLOCK_N, BLOCK_K, SPLIT_K, SWAP_AB, C_N_MAJOR, NUM_STAGES) \
else if (K == K_) dual_gemm_launch<K_, BLOCK_M, BLOCK_N, BLOCK_K, SPLIT_K, SWAP_AB, C_N_MAJOR, NUM_STAGES>(M, N, (const char*)A.data_ptr(), (const char*)B1.data_ptr(), (const char*)B2.data_ptr(), (const char*)SFA.data_ptr(), (const char*)SFB1.data_ptr(), (const char*)SFB2.data_ptr(), (half*)C.data_ptr(), (float*)buf.data_ptr());
if (false) {}
// Dual GEMM configs - use SWAP_AB=false and SPLIT_K=1 (split-K has bugs)
// submission-v7: Hybrid config - BLOCK_N=128 for K≥7168 (wins 43-44%), BLOCK_N=64 for K≤4096 (avoid regression)
LAUNCH(16384, 128, 128, 256, 1, false, true, 5) // Changed: NUM_STAGES 6→5 to fit SMEM
LAUNCH( 7168, 128, 128, 256, 1, false, true, 4) // Changed: BLOCK_N 64→128, NUM_STAGES 5→4
LAUNCH( 4096, 128, 64, 256, 1, false, false, 5) // Keep baseline: BLOCK_N=64, NUM_STAGES=5
LAUNCH( 2048, 128, 64, 256, 1, false, false, 5) // Keep baseline
// the rest - keep baseline config
LAUNCH( 256, 128, 64, 256, 1, false, false, 5)
LAUNCH( 512, 128, 64, 256, 1, false, false, 5)
LAUNCH(1536, 128, 64, 256, 1, false, false, 5)
LAUNCH(2304, 128, 64, 256, 1, false, false, 5)
else {
fprintf(stderr, "No matching kernel found for K=%d\n", K);
exit(1);
}
#undef LAUNCH
return C;
}
TORCH_LIBRARY(dual_gemm_module, m) {
m.def("dual_gemm_silu(Tensor A, Tensor B1, Tensor B2, Tensor SFA, Tensor SFB1, Tensor SFB2, Tensor(a!) C, Tensor(b!) buf) -> Tensor");
m.impl("dual_gemm_silu", &dual_gemm_silu);
}
"""
# Compile the inline CUDA module using TORCH_LIBRARY registration
load_inline(
name='dual_gemm_module_compile',
cpp_sources='',
cuda_sources=CUDA_SRC_UTILS + CUDA_SRC_KERNEL + CUDA_SRC_WRAPPER,
verbose=True,
is_python_module=False,
extra_cuda_cflags=[
'-O3',
'-gencode=arch=compute_100a,code=sm_100a',
'--use_fast_math',
'--expt-relaxed-constexpr',
'--relocatable-device-code=false',
'-lineinfo',
'-Xptxas=-v'
],
extra_ldflags=['-lcuda']
)
# Access the registered function through torch.ops
dual_gemm_op = torch.ops.dual_gemm_module.dual_gemm_silu
def custom_kernel(data):
a, b1, b2, sfa, sfb1, sfb2, sfa_permuted, sfb1_permuted, sfb2_permuted, c = data
M, K, L = a.shape
N = b1.shape[0]
# Handle L dimension by processing each batch slice
# Input tensors have shape [M, K, L] for a, [N, K, L] for b1/b2
# Scale factors have shape [M, K//16, L] and [N, K//16, L]
# Output c has shape [M, N, L]
# Allocate buffer for intermediate results (needed for split-K)
buf = torch.zeros((M, N), dtype=torch.float32, device='cuda')
# Process each L slice separately (reference implementation loops over L)
for l_idx in range(L):
# Extract 2D slices [M, K] and [N, K]
a_slice = a[:, :, l_idx]
b1_slice = b1[:, :, l_idx]
b2_slice = b2[:, :, l_idx]
sfa_slice = sfa_permuted[:, :, l_idx]
sfb1_slice = sfb1_permuted[:, :, l_idx]
sfb2_slice = sfb2_permuted[:, :, l_idx]
c_slice = c[:, :, l_idx]
# Call the kernel using Tensor objects (TORCH_LIBRARY approach)
# The C++ function extracts M, K, N from tensor shapes
dual_gemm_op(a_slice, b1_slice, b2_slice, sfa_slice, sfb1_slice, sfb2_slice, c_slice, buf)
return c
scrolls · 873 lines total
Source code from GPU Mode and the KernelBot dataset · June 9 Researcher Reciprocity License v1.0
Changes from previous submission
Against this author's previous submission submission 297615.
Best evidence level for this revision: reported
JSON