submission 370774
_spatters · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 686 lines, June 9 Researcher Reciprocity License v1.0.
v5.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-modal-nvfp4-dual-gemm-370774?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:15251ec9fb9e8224bc2834ec806f7bbfbc9a3af7852ffbcaea5e87bb85c9a022
license declaredunknown
license concludedunknown
authors_spatters
imported2026-08-26
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
async-copy
void cp_async_bulk_tensor_2d(const void *tensor_map_ptr, uint32_t dst_addr, int x_coord, int y_coord, uint32_t membar_addr) {fp4
PyTorch reference implementation of NVFP4 block-scaled GEMV.fp8
void cp_async_bulk(const __nv_fp8_e4m3 *src_addr, __nv_fp8_e4m3 *dst_addr, uint32_t cp_size, uint32_t membar_addr) {mbarrier
void mbarrier_wait(int mbar_addr, uint32_t phase) {shared-memory
extern __shared__ __align__(1024) char smem[];stages = 5
constexpr int NUM_STAGES = 5;tcgen05
"tcgen05.mma.cta_group::1.kind::mxf4nvf4.block_scale.block16 [%0], %1, %2, %3, [%4], [%5], p;\n\t"tma
CUtensorMap* tensor_map,vector-width = half2
reinterpret_cast<half2 *>(c_regs)[i] = __float22half2_rn(reinterpret_cast<float2*>(d1_regs)[i]);warp-specialization
asm volatile("setmaxnreg.dec.sync.aligned.u32 40;");Kernel source
v5.py686 lines
#!POPCORN leaderboard modal_nvfp4_dual_gemm
import os
os.environ["TORCH_CUDA_ARCH_LIST"] = "10.0"
import torch
from torch.utils.cpp_extension import load_inline
from task import input_t, output_t
cuda_source = r"""
#include<cstdio>
#include<cstdlib>
#include<stddef.h>
#include<cuda.h>
#include<cuda_fp4.h>
#include<cuda_fp16.h>
#define ceilDiv(x, y) (((x) + (y) - 1) / (y))
__forceinline__ __device__ float tanh_silu (float a, float b)
{
float r;
float a2 = 0.5*a;
float tmp = a2*b;
r = __tanhf(a2);
//asm ("tanh.approx.f32 %0,%1; \n\t" : "=f"(r) : "f"(a2));
return fmaf(tmp, r, tmp);
}
/* compute the sigmoid function 1/(1+exp(-x)) */
__forceinline__ __device__ float my_sigmoidf (float a)
{
#if USE_TANH
return fmaf (0.5, __tanhf (0.5f * a), 0.5f);
#else // USE_TANH
constexpr float L2E = 1.442695041f; // log2(exp(1))
float t, d, e, r;
t = -L2E * a;
asm ("ex2.approx.ftz.f32 %0,%1;\n\t" : "=f"(e) : "f"(t));
d = e + 1.0f;
asm ("rcp.approx.ftz.f32 %0,%1;\n\t" : "=f"(r) : "f"(d));
return r;
#endif // USE_TANH
}
__forceinline__ __device__ float silu (float a, float b)
{
#if USE_TANH
return fmaf (0.5, __tanhf (0.5f * a), 0.5f);
#else // USE_TANH
constexpr float L2E = 1.442695041f; // log2(exp(1))
float t, d, e, r, tmp;
t = -L2E * a;
tmp = a * b;
asm ("ex2.approx.ftz.f32 %0,%1;\n\t" : "=f"(e) : "f"(t));
d = e + 1.0f;
asm ("rcp.approx.ftz.f32 %0,%1;\n\t" : "=f"(r) : "f"(d));
return r*tmp;
#endif // USE_TANH
}
inline const char* cuGetErrorNameSafe(CUresult err) {
const char* name = nullptr;
// cuGetErrorName returns CUDA_SUCCESS even if it can't map the error
if (cuGetErrorName(err, &name) != CUDA_SUCCESS || !name) return "<unknown>";
return name;
}
inline const char* cuGetErrorStringSafe(CUresult err) {
const char* str = nullptr;
if (cuGetErrorString(err, &str) != CUDA_SUCCESS || !str) return "<no description>";
return str;
}
inline void cuCheckImpl(CUresult err, const char* expr, const char* file, int line) {
if (err != CUDA_SUCCESS) {
std::fprintf(stderr,
"CUDA Driver API error: %s returned %d (%s): %s at %s:%d\n",
expr,
(int)err,
cuGetErrorNameSafe(err),
cuGetErrorStringSafe(err),
file, line
);
std::fflush(stderr);
std::abort();
}
}
// Elect one thread in the warp. The elected thread gets its predicate set to true, all others obtain false.
__device__
uint32_t elect_one_sync()
{
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;
}
__device__
void mbarrier_wait(int mbar_addr, uint32_t phase) {
uint32_t ticks = 0x989680;
asm volatile(
"{\n\t"
".reg .pred P1; \n\t"
"LAB_WAIT: \n\t"
"mbarrier.try_wait.parity.shared::cta.b64 P1, [%0], %1, %2; \n\t"
"@P1 bra DONE; \n\t"
"bra LAB_WAIT; \n\t"
"DONE: \n\t"
"}"
:
: "r"(mbar_addr), "r"(phase), "r"(ticks)
: "memory");
}
#define CU_CHECK(expr) cuCheckImpl((expr), #expr, __FILE__, __LINE__)
#define CU_CHECK_RESULT_WHAT(err, what) cuCheckImpl((err), (what), __FILE__, __LINE__)
void create_tensor_map(
CUtensorMap* tensor_map,
__nv_fp4_e2m1* globalPtr,
uint64_t global_M,
uint64_t global_K,
uint32_t block_M,
uint32_t block_K
) {
constexpr uint32_t rank = 2;
uint64_t size[rank] = {global_K, global_M};
uint64_t stride[rank - 1] = {global_K / 2 };
uint32_t box_size[rank] = {block_K, block_M};
uint32_t elem_stride[rank] = {1, 1};
CUtensorMapSwizzle swizzle_mode;
if (block_K==256) {
swizzle_mode = CUtensorMapSwizzle::CU_TENSOR_MAP_SWIZZLE_128B;
} else if (block_K==128) {
swizzle_mode = CUtensorMapSwizzle::CU_TENSOR_MAP_SWIZZLE_64B;
} else if (block_K==64) {
swizzle_mode = CUtensorMapSwizzle::CU_TENSOR_MAP_SWIZZLE_32B;
}
CUresult res = cuTensorMapEncodeTiled(
tensor_map, // CUtensorMap *tensorMap,
CUtensorMapDataType::CU_TENSOR_MAP_DATA_TYPE_16U4_ALIGN8B,
rank, // cuuint32_t tensorRank,
(void *)globalPtr, // void *globalAddress,
size, // const cuuint64_t *globalDim,
stride, // const cuuint64_t *globalStrides,
box_size, // const cuuint32_t *boxDim,
elem_stride, // const cuuint32_t *elementStrides,
CUtensorMapInterleave::CU_TENSOR_MAP_INTERLEAVE_NONE,
swizzle_mode,
CUtensorMapL2promotion::CU_TENSOR_MAP_L2_PROMOTION_NONE,
CUtensorMapFloatOOBfill::CU_TENSOR_MAP_FLOAT_OOB_FILL_NONE
);
CU_CHECK_RESULT_WHAT(res, "cuTensorMapEncodeTiled");
}
template<int TILE_SIZE>
__device__ __forceinline__
void get_tile(int idx, int& tile_id, int& offset) {
static_assert((TILE_SIZE & (TILE_SIZE - 1)) == 0, "Must be power of 2");
constexpr int mask = TILE_SIZE - 1;
constexpr int shift = __builtin_ctz(TILE_SIZE);
tile_id = idx >> shift;
offset = idx & mask;
}
__device__ __forceinline__
void cp_async_bulk_tensor_2d(const void *tensor_map_ptr, uint32_t dst_addr, int x_coord, int y_coord, uint32_t membar_addr) {
asm volatile("cp.async.bulk.tensor.2d.shared::cta.global.mbarrier::complete_tx::bytes [%0], [%1, {%2, %3}], [%4];"
:: "r"(dst_addr), "l"(tensor_map_ptr), "r"(x_coord), "r"(y_coord), "r"(membar_addr) : "memory");
}
__device__ __forceinline__
void cp_async_bulk(const __nv_fp8_e4m3 *src_addr, __nv_fp8_e4m3 *dst_addr, uint32_t cp_size, uint32_t membar_addr) {
uint32_t ptx_dst_addr = __cvta_generic_to_shared(dst_addr);
asm volatile("cp.async.bulk.shared::cta.global.mbarrier::complete_tx::bytes [%0], [%1], %2, [%3];"
:: "r"(ptx_dst_addr), "l"(src_addr), "r"(cp_size), "r"(membar_addr) : "memory");
}
__device__ __forceinline__
void mbarrier_init(uint32_t membar_addr, int count) {
asm volatile("mbarrier.init.shared::cta.b64 [%0], %1;" :: "r"(membar_addr), "r"(count));
asm volatile("fence.proxy.async.shared::cta;");
}
__device__ __forceinline__
void mbarrier_arrive_expect_tx(uint32_t membar_addr, int tx_count) {
asm volatile("mbarrier.arrive.expect_tx.release.cta.shared::cta.b64 _, [%0], %1;" :: "r"(membar_addr), "r"(tx_count) : "memory");
}
__device__ __forceinline__
void tcgen05_mma(uint32_t output_addr, uint64_t a_desc, uint64_t b_desc, uint32_t sfa_tmem_addr, uint32_t sfb_tmem_addr, uint32_t i_desc, uint32_t enable_input_d) {
asm volatile(
"{\n\t"
".reg .pred p;\n\t"
"setp.ne.u32 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\t"
:: "r"(output_addr), "l"(a_desc), "l"(b_desc), "r"(i_desc),
"r"(sfa_tmem_addr), "r"(sfb_tmem_addr), "r"(enable_input_d));
}
__device__ __forceinline__
void tcgen05_ld_32(float* D, uint32_t tmem_addr) {
asm volatile(
"{\n\t"
"tcgen05.ld.sync.aligned.32x32b.x32.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];"
"}\n\t"
: "=f"(D[0]), "=f"(D[1]), "=f"(D[2]), "=f"(D[3]), "=f"(D[4]), "=f"(D[5]), "=f"(D[6]), "=f"(D[7]),
"=f"(D[8]), "=f"(D[9]), "=f"(D[10]), "=f"(D[11]), "=f"(D[12]), "=f"(D[13]), "=f"(D[14]), "=f"(D[15]),
"=f"(D[16]), "=f"(D[17]), "=f"(D[18]), "=f"(D[19]), "=f"(D[20]), "=f"(D[21]), "=f"(D[22]), "=f"(D[23]),
"=f"(D[24]), "=f"(D[25]), "=f"(D[26]), "=f"(D[27]), "=f"(D[28]), "=f"(D[29]), "=f"(D[30]), "=f"(D[31])
:
"r"(tmem_addr)
);
}
__device__ __forceinline__
void tcgen05_ld_64(float* D, uint32_t tmem_addr) {
asm volatile(
"{\n\t"
"tcgen05.ld.sync.aligned.32x32b.x64.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];"
"}\n\t"
: "=f"(D[0]), "=f"(D[1]), "=f"(D[2]), "=f"(D[3]), "=f"(D[4]), "=f"(D[5]), "=f"(D[6]), "=f"(D[7]),
"=f"(D[8]), "=f"(D[9]), "=f"(D[10]), "=f"(D[11]), "=f"(D[12]), "=f"(D[13]), "=f"(D[14]), "=f"(D[15]),
"=f"(D[16]), "=f"(D[17]), "=f"(D[18]), "=f"(D[19]), "=f"(D[20]), "=f"(D[21]), "=f"(D[22]), "=f"(D[23]),
"=f"(D[24]), "=f"(D[25]), "=f"(D[26]), "=f"(D[27]), "=f"(D[28]), "=f"(D[29]), "=f"(D[30]), "=f"(D[31]),
"=f"(D[32]), "=f"(D[33]), "=f"(D[34]), "=f"(D[35]), "=f"(D[36]), "=f"(D[37]), "=f"(D[38]), "=f"(D[39]),
"=f"(D[40]), "=f"(D[41]), "=f"(D[42]), "=f"(D[43]), "=f"(D[44]), "=f"(D[45]), "=f"(D[46]), "=f"(D[47]),
"=f"(D[48]), "=f"(D[49]), "=f"(D[50]), "=f"(D[51]), "=f"(D[52]), "=f"(D[53]), "=f"(D[54]), "=f"(D[55]),
"=f"(D[56]), "=f"(D[57]), "=f"(D[58]), "=f"(D[59]), "=f"(D[60]), "=f"(D[61]), "=f"(D[62]), "=f"(D[63])
:
"r"(tmem_addr)
);
}
// gau
__device__ inline
constexpr uint64_t desc_encode(uint64_t x) { return (x & 0x3FFFFULL) >> 4ULL; };
template<int M_BLOCK, int N_BLOCK, int K_BLOCK>
__global__ void gemm_kernel2(
const __grid_constant__ CUtensorMap tensor_map_A,
const __grid_constant__ CUtensorMap tensor_map_B,
const __nv_fp8_e4m3* SFA,
const __nv_fp8_e4m3* SFB,
half* C,
int M,
int N,
int K
) {
int global_tidx = blockIdx.x*128 + blockIdx.y*128 + threadIdx.x;
}
template<int M_BLOCK, int N_BLOCK, int K_BLOCK, int K_MMA, int NUM_STAGES>
__global__ void dual_gemm_kernel(
const __grid_constant__ CUtensorMap tensor_map_A,
const __grid_constant__ CUtensorMap tensor_map_B1,
const __grid_constant__ CUtensorMap tensor_map_B2,
const __nv_fp8_e4m3* SFA,
const __nv_fp8_e4m3* SFB1,
const __nv_fp8_e4m3* SFB2,
half* C,
int M,
int N,
int K
) {
extern __shared__ __align__(1024) char smem[];
constexpr uint32_t A_SIZE = M_BLOCK*K_BLOCK/2;
constexpr uint32_t B_SIZE = N_BLOCK*K_BLOCK/2;
constexpr uint32_t SFA_SIZE = M_BLOCK*K_BLOCK/16;
constexpr uint32_t SFB_SIZE = M_BLOCK*K_BLOCK/16;
__nv_fp4x2_e2m1* As = reinterpret_cast<__nv_fp4x2_e2m1*>(smem);
__nv_fp4x2_e2m1* B1s = As + NUM_STAGES*A_SIZE;
__nv_fp4x2_e2m1* B2s = B1s + NUM_STAGES*B_SIZE;
__nv_fp8_e4m3* sfas = reinterpret_cast<__nv_fp8_e4m3*>(B2s + NUM_STAGES*1*B_SIZE);
__nv_fp8_e4m3* sfb1s = sfas + NUM_STAGES*SFA_SIZE;
__nv_fp8_e4m3* sfb2s = sfb1s + NUM_STAGES*SFB_SIZE;
uint64_t* tma_mbar = reinterpret_cast<uint64_t*>(sfb2s + NUM_STAGES*1*SFB_SIZE);
uint64_t* mma_mbar = tma_mbar + NUM_STAGES;
uint64_t* main_loop_mbar = mma_mbar + NUM_STAGES;
uint32_t* tmem_addr = reinterpret_cast<uint32_t*>(main_loop_mbar + 1);
//__shared__ uint32_t tmem_addr;
//__shared__ uint32_t tmem_addr_ptr;
uint32_t tmem_addr_ptr = __cvta_generic_to_shared(tmem_addr);
int threadID = threadIdx.x;
int warpID, laneID;
get_tile<32>(threadID, warpID, laneID);
int warpID_4 = warpID%4;
if (warpID < 4) {
asm volatile("setmaxnreg.dec.sync.aligned.u32 40;");
} else {
asm volatile("setmaxnreg.dec.sync.aligned.u32 96;");
}
uint32_t tma_mbar_addr = __cvta_generic_to_shared(tma_mbar);
uint32_t mma_mbar_addr = __cvta_generic_to_shared(mma_mbar);
uint32_t main_loop_mbar_addr = __cvta_generic_to_shared(main_loop_mbar);
// per block we load M/N_BLOCK * K/16 sfa/b scale factors
uint32_t block_sfa_offset = blockIdx.x * (M_BLOCK*K/16);
uint32_t block_sfb_offset = (blockIdx.y/2) * (M_BLOCK*K/16);
uint32_t block_m_idx = blockIdx.x * M_BLOCK;
uint32_t block_n_idx = blockIdx.y * N_BLOCK;
uint32_t global_m_id = block_m_idx + 32*warpID_4 + laneID;
uint32_t global_n_id = block_n_idx;
int block_idx = blockIdx.y*blockDim.x + blockIdx.x;
int global_thread_idx = block_idx*128 + threadID;
constexpr uint32_t tmem_ncols = N_BLOCK * 4;
float d1_regs[64];
float d2_regs[64];
half c_regs[64];
//constexpr uint32_t cp_size = sizeof(As) + sizeof(Bs) + sizeof(sfas) + 0*sizeof(sfbs);
constexpr uint32_t cp_size = A_SIZE + 2*B_SIZE + SFA_SIZE + 2*SFB_SIZE;
constexpr uint32_t ins_desc = (1U << 7U) | (1U << 10U) | (8U << 17U) | (1U << 27U);
// gau nerst
// set up shared memory descriptors for A and B
// https://docs.nvidia.com/cuda/parallel-thread-execution/#tcgen05-shared-memory-descriptor
// 128-byte swizzling. LBO is implied to be 1.
// 128B => 2
// 64B => 4
// 32B => 6
//const uint64_t swizzle_enum = 4;
constexpr uint64_t swizzle_enum =
(K_BLOCK==256) ? 2 :
(K_BLOCK==128) ? 4 :
6;
auto make_desc_AB = [](uint32_t addr) -> uint64_t {
constexpr int SBO = 8 * K_BLOCK/2;
return desc_encode(addr) | (desc_encode(SBO) << 32ULL) | (1ULL << 46ULL) | (swizzle_enum << 61ULL);
};
// no swizzling
auto make_desc_SF = [](uint32_t addr) -> uint64_t {
constexpr int SBO = 8 * 16;
return desc_encode(addr) | (desc_encode(SBO) << 32ULL) | (1ULL << 46ULL);
};
// launch tma copy
if (warpID == 0 && elect_one_sync()) {
#pragma unroll
for (int n_stage=0; n_stage<NUM_STAGES; ++n_stage) {
mbarrier_init(tma_mbar_addr + n_stage*8, 1);
mbarrier_init(mma_mbar_addr + n_stage*8, 1);
}
mbarrier_init(main_loop_mbar_addr, 1);
}
if (warpID == 1) {
// allocate tensor memory: all 128 lanes(rows) are allocated
// number of columns must be in {32, 64, 128, 256, 512}
// need to allocate for mma output and for sfa/sfb (could we do these in separate allocations / would that use less memory)?
asm volatile("tcgen05.alloc.cta_group::1.sync.aligned.shared::cta.b32 [%0], %1;" :: "r"(tmem_addr_ptr), "r"(tmem_ncols): "memory");
}
constexpr int NUM_MMAs = K_BLOCK / K_MMA;
constexpr uint32_t c1_tmem_addr = 0;
constexpr uint32_t c2_tmem_addr = N_BLOCK;
constexpr uint32_t sfa_tmem_addr = 2 * N_BLOCK;
constexpr uint32_t sfb1_tmem_addr = sfa_tmem_addr + 16;
constexpr uint32_t sfb2_tmem_addr = sfb1_tmem_addr + 16;
uint64_t a_smem_dsec = make_desc_AB(__cvta_generic_to_shared(As));
uint64_t b1_smem_dsec = make_desc_AB(__cvta_generic_to_shared(B1s));
uint64_t b2_smem_dsec = make_desc_AB(__cvta_generic_to_shared(B2s));
uint64_t sfa_smem_dsec = make_desc_SF(__cvta_generic_to_shared(sfas));
uint64_t sfb1_smem_dsec = make_desc_SF(__cvta_generic_to_shared(sfb1s));
uint64_t sfb2_smem_dsec = make_desc_SF(__cvta_generic_to_shared(sfb2s));
__syncthreads();
int tma_phase[NUM_STAGES] = {0};
int mma_phase[NUM_STAGES] = {0};
if (warpID==0 && elect_one_sync()) {
// prefill the shared memory buffers before starting mmas
for (int n_stage=0; n_stage<NUM_STAGES; ++n_stage) {
uint32_t stage_tma_mbar_addr = tma_mbar_addr + n_stage*8;
cp_async_bulk_tensor_2d(&tensor_map_A, __cvta_generic_to_shared(As) + n_stage*A_SIZE, n_stage*K_BLOCK, block_m_idx, stage_tma_mbar_addr);
cp_async_bulk_tensor_2d(&tensor_map_B1, __cvta_generic_to_shared(B1s) + n_stage*B_SIZE, n_stage*K_BLOCK, block_n_idx, stage_tma_mbar_addr);
cp_async_bulk_tensor_2d(&tensor_map_B2, __cvta_generic_to_shared(B2s) + n_stage*B_SIZE, n_stage*K_BLOCK, block_n_idx, stage_tma_mbar_addr);
cp_async_bulk(SFA + block_sfa_offset + (n_stage*K_BLOCK/16)*M_BLOCK, sfas + n_stage*SFA_SIZE, SFA_SIZE, stage_tma_mbar_addr);
cp_async_bulk(SFB1 + block_sfb_offset + (n_stage*K_BLOCK/16)*M_BLOCK, sfb1s + n_stage*SFB_SIZE, SFB_SIZE, stage_tma_mbar_addr);
cp_async_bulk(SFB2 + block_sfb_offset + (n_stage*K_BLOCK/16)*M_BLOCK, sfb2s + n_stage*SFB_SIZE, SFB_SIZE, stage_tma_mbar_addr);
mbarrier_arrive_expect_tx(stage_tma_mbar_addr, cp_size);
}
for (int k_idx=NUM_STAGES*K_BLOCK; k_idx<K; k_idx+=K_BLOCK) {
int n_stage = (k_idx/K_BLOCK) % NUM_STAGES;
uint32_t stage_tma_mbar_addr = tma_mbar_addr + n_stage*8;
uint32_t stage_mma_mbar_addr = mma_mbar_addr + n_stage*8;
mbarrier_wait(stage_mma_mbar_addr, mma_phase[n_stage]);
mma_phase[n_stage] ^= 1;
cp_async_bulk_tensor_2d(&tensor_map_A, __cvta_generic_to_shared(As) + n_stage*A_SIZE, k_idx, block_m_idx, stage_tma_mbar_addr);
cp_async_bulk_tensor_2d(&tensor_map_B1, __cvta_generic_to_shared(B1s) + n_stage*B_SIZE, k_idx, block_n_idx, stage_tma_mbar_addr);
cp_async_bulk_tensor_2d(&tensor_map_B2, __cvta_generic_to_shared(B2s) + n_stage*B_SIZE, k_idx, block_n_idx, stage_tma_mbar_addr);
cp_async_bulk(SFA + block_sfa_offset + (k_idx/16)*M_BLOCK, sfas + n_stage*SFA_SIZE, SFA_SIZE, stage_tma_mbar_addr);
cp_async_bulk(SFB1 + block_sfb_offset + (k_idx/16)*M_BLOCK, sfb1s + n_stage*SFB_SIZE, SFB_SIZE, stage_tma_mbar_addr);
cp_async_bulk(SFB2 + block_sfb_offset + (k_idx/16)*M_BLOCK, sfb2s + n_stage*SFB_SIZE, SFB_SIZE, stage_tma_mbar_addr);
mbarrier_arrive_expect_tx(stage_tma_mbar_addr, cp_size);
}
}
else if (warpID==1 && elect_one_sync()) {
for (int k_idx=0; k_idx<K; k_idx+=K_BLOCK) {
int n_stage = (k_idx/K_BLOCK) % NUM_STAGES;
uint32_t stage_A = __cvta_generic_to_shared(As) + n_stage*A_SIZE;
uint32_t stage_B1 = __cvta_generic_to_shared(B1s) + n_stage*B_SIZE;
uint32_t stage_B2 = __cvta_generic_to_shared(B2s) + n_stage*B_SIZE;
uint32_t stage_SFA = __cvta_generic_to_shared(sfas) + n_stage*SFA_SIZE;
uint32_t stage_SFB1 = __cvta_generic_to_shared(sfb1s) + n_stage*SFB_SIZE;
uint32_t stage_SFB2 = __cvta_generic_to_shared(sfb2s) + n_stage*SFB_SIZE;
uint32_t stage_tma_mbar_addr = tma_mbar_addr + n_stage*8;
uint32_t stage_mma_mbar_addr = mma_mbar_addr + n_stage*8;
mbarrier_wait(stage_tma_mbar_addr, tma_phase[n_stage]);
tma_phase[n_stage] ^=1 ;
for (int k_inner=0; k_inner<NUM_MMAs; k_inner++) {
sfa_smem_dsec = make_desc_SF(stage_SFA + k_inner*512);
sfb1_smem_dsec = make_desc_SF(stage_SFB1 + k_inner*512);
sfb2_smem_dsec = make_desc_SF(stage_SFB2 + k_inner*512);
asm volatile("tcgen05.cp.cta_group::1.32x128b.warpx4 [%0], %1;" :: "r"(sfa_tmem_addr+4*k_inner), "l"(sfa_smem_dsec));
asm volatile("tcgen05.cp.cta_group::1.32x128b.warpx4 [%0], %1;" :: "r"(sfb1_tmem_addr+4*k_inner), "l"(sfb1_smem_dsec));
asm volatile("tcgen05.cp.cta_group::1.32x128b.warpx4 [%0], %1;" :: "r"(sfb2_tmem_addr+4*k_inner), "l"(sfb2_smem_dsec));
}
// launch mma
for (int k_inner=0; k_inner<NUM_MMAs; k_inner++) {
a_smem_dsec = make_desc_AB(stage_A + k_inner*32);
b1_smem_dsec = make_desc_AB(stage_B1 + k_inner*32);
b2_smem_dsec = make_desc_AB(stage_B2 + k_inner*32);
tcgen05_mma(c1_tmem_addr, a_smem_dsec, b1_smem_dsec, sfa_tmem_addr + 4*k_inner, sfb1_tmem_addr + 4*k_inner + 2*(blockIdx.y%2), ins_desc, k_idx+k_inner*64);
tcgen05_mma(c2_tmem_addr, a_smem_dsec, b2_smem_dsec, sfa_tmem_addr + 4*k_inner, sfb2_tmem_addr + 4*k_inner + 2*(blockIdx.y%2), ins_desc, k_idx+k_inner*64);
}
asm volatile("tcgen05.commit.cta_group::1.mbarrier::arrive::one.b64 [%0];" :: "r"(stage_mma_mbar_addr));
}
// signal mainloop done
asm volatile("tcgen05.commit.cta_group::1.mbarrier::arrive::one.shared::cluster.b64 [%0];"
:: "r"(main_loop_mbar_addr) : "memory");
}
else if (warpID>3) {
// epilogue and store
mbarrier_wait(main_loop_mbar_addr, 0);
asm volatile("tcgen05.fence::after_thread_sync;");
tcgen05_ld_64(d1_regs, ((warpID_4*32) << 16));
tcgen05_ld_64(d2_regs, ((warpID_4*32) << 16) + N_BLOCK);
asm volatile("tcgen05.wait::ld.sync.aligned;");
#pragma unroll
for (int i=0; i<64; ++i) {
//d1_regs[i] = my_sigmoidf(d1_regs[i]) * d1_regs[i];
d1_regs[i] = tanh_silu(d1_regs[i], d2_regs[i]);
}
/*
#pragma unroll
for (int i=0; i<64; ++i) {
d1_regs[i] *= d2_regs[i];
}
*/
#pragma unroll
for (int i=0; i<32; ++i) {
reinterpret_cast<half2 *>(c_regs)[i] = __float22half2_rn(reinterpret_cast<float2*>(d1_regs)[i]);
}
// store to C
for (int j=0; j<64; ++j) {
if ((global_m_id < M) && (global_n_id + j < N)) {
C[(global_n_id+j)*M + global_m_id] = c_regs[j];
//C[(global_m_id)*N + global_n_id+j] = c_regs[j];
}
}
}
__syncthreads();
if (warpID == 0) // deallocate tmem. tmem address should be 0.
asm volatile("tcgen05.dealloc.cta_group::1.sync.aligned.b32 %0, %1;" :: "r"(*tmem_addr), "r"(tmem_ncols));
}
template<int M_BLOCK, int N_BLOCK, int K_BLOCK, int K_MMA, int NUM_STAGES>
void launch_dual_gemm(
const CUtensorMap tensor_map_A,
const CUtensorMap tensor_map_B1,
const CUtensorMap tensor_map_B2,
const __nv_fp8_e4m3* SFA,
const __nv_fp8_e4m3* SFB1,
const __nv_fp8_e4m3* SFB2,
half* C,
int M,
int N,
int K,
int L)
{
int threads = 256;
dim3 grid(ceilDiv(M, M_BLOCK), ceilDiv(N, N_BLOCK), L);
constexpr size_t SMEM_BYTES = NUM_STAGES * (M_BLOCK*K_BLOCK/2 + 2 * N_BLOCK*K_BLOCK/2 + 3*M_BLOCK*(K_BLOCK/16) + 4*sizeof(uint64_t));
//printf("SMEM_BYTES: %'d\n", SMEM_BYTES);
auto kernel_fn = dual_gemm_kernel<M_BLOCK, N_BLOCK, K_BLOCK, K_MMA, NUM_STAGES>;
cudaFuncSetAttribute(kernel_fn, cudaFuncAttributeMaxDynamicSharedMemorySize, SMEM_BYTES);
kernel_fn<<<grid, threads, SMEM_BYTES>>>(tensor_map_A, tensor_map_B1, tensor_map_B2, SFA, SFB1, SFB2, C, M, N, K);
}
torch::Tensor dual_gemm_cuda(torch::Tensor A, torch::Tensor B1, torch::Tensor B2,
torch::Tensor SFA, torch::Tensor SFB1, torch::Tensor SFB2, torch::Tensor C) {
TORCH_CHECK(A.device().is_cuda(), "Tensor A must be a CUDA tensor");
TORCH_CHECK(B1.device().is_cuda(), "Tensor B must be a CUDA tensor");
TORCH_CHECK(B2.device().is_cuda(), "Tensor B must be a CUDA tensor");
TORCH_CHECK(SFA.device().is_cuda(), "Tensor SFA must be a CUDA tensor");
TORCH_CHECK(SFB1.device().is_cuda(), "Tensor SFB must be a CUDA tensor");
TORCH_CHECK(SFB2.device().is_cuda(), "Tensor SFB must be a CUDA tensor");
TORCH_CHECK(C.device().is_cuda(), "Tensor C must be a CUDA tensor");
int M = A.size(0);
int N = B1.size(0);
int K = A.size(1) * 2; // A is type nv_fp4x2
int L = A.size(2);
//printf("Problem size M: %d, K: %d, N: %d, L: %d \n", M, K, N, L);
auto A_ptr = reinterpret_cast<__nv_fp4_e2m1*>(A.data_ptr());
auto B1_ptr = reinterpret_cast<__nv_fp4_e2m1*>(B1.data_ptr());
auto B2_ptr = reinterpret_cast<__nv_fp4_e2m1*>(B2.data_ptr());
auto SFA_ptr = reinterpret_cast<__nv_fp8_e4m3*>(SFA.data_ptr());
auto SFB1_ptr = reinterpret_cast<__nv_fp8_e4m3*>(SFB1.data_ptr());
auto SFB2_ptr = reinterpret_cast<__nv_fp8_e4m3*>(SFB2.data_ptr());
auto C_ptr = reinterpret_cast<__half*>(C.data_ptr());
CUtensorMap tensor_map_A{};
CUtensorMap tensor_map_B1{};
CUtensorMap tensor_map_B2{};
constexpr int M_BLOCK = 128;
constexpr int N_BLOCK = 64;
constexpr int K_BLOCK = 256;
constexpr int K_MMA = 64;
constexpr int NUM_STAGES = 5;
create_tensor_map(&tensor_map_A, A_ptr, M, K, (uint32_t)M_BLOCK, (uint32_t)K_BLOCK);
create_tensor_map(&tensor_map_B1, B1_ptr, N, K, (uint32_t)N_BLOCK, (uint32_t)K_BLOCK);
create_tensor_map(&tensor_map_B2, B2_ptr, N, K, (uint32_t)N_BLOCK, (uint32_t)K_BLOCK);
launch_dual_gemm<M_BLOCK, N_BLOCK, K_BLOCK, K_MMA, NUM_STAGES>(tensor_map_A, tensor_map_B1, tensor_map_B2, SFA_ptr, SFB1_ptr, SFB2_ptr, C_ptr, M, N, K, L);
cudaError_t err = cudaGetLastError();
if (err != cudaSuccess) {
throw std::runtime_error(cudaGetErrorString(err));
}
auto C_transpose = C.view({N, M, 1}).transpose(0, 1);
//auto C_transpose2 = C.transpose(0, 1);
/*
std::cout << "C shape/strides: " << C.sizes() << ", " << C.strides() << std::endl;
std::cout << "C_transpose shape/strides: " << C_transpose.sizes() << ", " << C_transpose.strides() << std::endl;
std::cout << "C_transpose2 shape/strides: " << C_transpose2.sizes() << ", " << C_transpose2.strides() << std::endl;
std::cout << "C[10,15]: " << C.index({10,15,0}) << std::endl;
std::cout << "C_transpose[10,15]: " << C_transpose.index({10,15,0}) << std::endl;
*/
return C_transpose;
}
"""
cpp_source = """
#include <torch/extension.h>
torch::Tensor dual_gemm_cuda(
torch::Tensor A,
torch::Tensor B1,
torch::Tensor B2,
torch::Tensor SFA,
torch::Tensor SFB1,
torch::Tensor SFB2,
torch::Tensor C);
"""
extra_cuda_cflags = [
"-O3",
"--use_fast_math",
"--fmad=true",
"--ftz=true",
#"-Xcompiler", "-fno-strict-aliasing",
# Aggressive math optimizations
"-Xptxas=-O3",
# Cache behavior
#"-Xptxas=-dlcm=ca",
# For debugging performance
"-Xptxas=--warn-on-spills",
"-Xptxas=-v",
# Blackwell target
#"--gpu-architecture=sm_100a",
"-gencode=arch=compute_100a,code=sm_100a",
]
extra_cflags = [
"-O3",
"-ffast-math",
"-fno-strict-aliasing",
]
dual_gemm_module = load_inline(
name='dual_gemm_cuda',
cpp_sources=cpp_source,
cuda_sources=cuda_source,
functions=['dual_gemm_cuda'],
verbose=True,
extra_cuda_cflags=extra_cuda_cflags,
extra_cflags=extra_cflags,
extra_ldflags=["-lcuda"],
)
def dual_gemm_cuda(A, B1, B2, SFA, SFB1, SFB2, C):
if not A.is_cuda or not B1.is_cuda or not SFA.is_cuda or not SFB1.is_cuda or not C.is_cuda:
raise RuntimeError("Both tensors must be on GPU")
return dual_gemm_module.dual_gemm_cuda(A, B1, B2, SFA, SFB1, SFB2, C)
# Helper function for ceiling division
def ceil_div(a, b):
return (a + b - 1) // b
def custom_kernel(
data: input_t,
) -> output_t:
"""
PyTorch reference implementation of NVFP4 block-scaled GEMV.
"""
a, b1, b2, sfa, sfb1, sfb2, sfa_permuted, sfb1_permuted, sfb2_permuted, c = data
m, k, l = a.shape
n, k, l = b1.shape
"""
print(f"K is {k}, n is {n}")
print(f"A shape {a_ref.shape}")
print(f"A shape {a_ref.stride()}")
print(f"SFA shape {sfa.shape}")
print(f"SFA shape {sfa.stride()}")
print(f"B1 shape {b1.shape}")
print(f"B1 strid {b2.stride()}")
print(f"SFB shape {sfb.shape}")
print(f"SFB shape {sfb.stride()}")
print(f"C shape {c_ref.shape}")
print(f"C shape {c_ref.stride()}")
"""
# Get dimensions from MxNxL layout
_, _, l = c.shape
#print(sfa.shape, sfa.stride())
#print(f"SFA[0,0:32,0]: {sfa[0,:32,0].reshape(-1,2)}")
out = dual_gemm_cuda(a, b1, b2, sfa_permuted, sfb1_permuted, sfb2_permuted, c)
#torch.cuda.synchronize()
#print(c_ref)
return out
scrolls · 686 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 319574.
- #!POPCORN leaderboard nvfp4_dual_gemm+ #!POPCORN leaderboard modal_nvfp4_dual_gemmimport osos.environ["TORCH_CUDA_ARCH_LIST"] = "10.0"⋯ 11 unchanged lines#define ceilDiv(x, y) (((x) + (y) - 1) / (y))++ __forceinline__ __device__ float tanh_silu (float a, float b)+ {+ float r;+ float a2 = 0.5*a;+ float tmp = a2*b;+ r = __tanhf(a2);+ //asm ("tanh.approx.f32 %0,%1; \n\t" : "=f"(r) : "f"(a2));+ return fmaf(tmp, r, tmp);+ }+/* compute the sigmoid function 1/(1+exp(-x)) */__forceinline__ __device__ float my_sigmoidf (float a){⋯ 10 unchanged lines#endif // USE_TANH}+ __forceinline__ __device__ float silu (float a, float b)+ {+ #if USE_TANH+ return fmaf (0.5, __tanhf (0.5f * a), 0.5f);+ #else // USE_TANH+ constexpr float L2E = 1.442695041f; // log2(exp(1))+ float t, d, e, r, tmp;+ t = -L2E * a;+ tmp = a * b;+ asm ("ex2.approx.ftz.f32 %0,%1;\n\t" : "=f"(e) : "f"(t));+ d = e + 1.0f;+ asm ("rcp.approx.ftz.f32 %0,%1;\n\t" : "=f"(r) : "f"(d));+ return r*tmp;+ #endif // USE_TANH+ }+inline const char* cuGetErrorNameSafe(CUresult err) {const char* name = nullptr;// cuGetErrorName returns CUDA_SUCCESS even if it can't map the error⋯ 74 unchanged linesuint64_t stride[rank - 1] = {global_K / 2 };uint32_t box_size[rank] = {block_K, block_M};uint32_t elem_stride[rank] = {1, 1};- CUtensorMapSwizzle swizzle_mode = (block_K==256) ? CUtensorMapSwizzle::CU_TENSOR_MAP_SWIZZLE_128B : CUtensorMapSwizzle::CU_TENSOR_MAP_SWIZZLE_32B;+ CUtensorMapSwizzle swizzle_mode;+ if (block_K==256) {+ swizzle_mode = CUtensorMapSwizzle::CU_TENSOR_MAP_SWIZZLE_128B;+ } else if (block_K==128) {+ swizzle_mode = CUtensorMapSwizzle::CU_TENSOR_MAP_SWIZZLE_64B;+ } else if (block_K==64) {+ swizzle_mode = CUtensorMapSwizzle::CU_TENSOR_MAP_SWIZZLE_32B;+ }+CUresult res = cuTensorMapEncodeTiled(tensor_map, // CUtensorMap *tensorMap,CUtensorMapDataType::CU_TENSOR_MAP_DATA_TYPE_16U4_ALIGN8B,⋯ 4 unchanged linesbox_size, // const cuuint32_t *boxDim,elem_stride, // const cuuint32_t *elementStrides,CUtensorMapInterleave::CU_TENSOR_MAP_INTERLEAVE_NONE,- //CUtensorMapSwizzle::CU_TENSOR_MAP_SWIZZLE_32B,- //CUtensorMapSwizzle::CU_TENSOR_MAP_SWIZZLE_128B,swizzle_mode,CUtensorMapL2promotion::CU_TENSOR_MAP_L2_PROMOTION_NONE,CUtensorMapFloatOOBfill::CU_TENSOR_MAP_FLOAT_OOB_FILL_NONE⋯ 177 unchanged linesconstexpr uint32_t tmem_ncols = N_BLOCK * 4;- float d1_regs[32];- float d2_regs[32];- half c_regs[32];+ float d1_regs[64];+ float d2_regs[64];+ half c_regs[64];//constexpr uint32_t cp_size = sizeof(As) + sizeof(Bs) + sizeof(sfas) + 0*sizeof(sfbs);constexpr uint32_t cp_size = A_SIZE + 2*B_SIZE + SFA_SIZE + 2*SFB_SIZE;⋯ 5 unchanged lines// set up shared memory descriptors for A and B// https://docs.nvidia.com/cuda/parallel-thread-execution/#tcgen05-shared-memory-descriptor// 128-byte swizzling. LBO is implied to be 1.+ // 128B => 2+ // 64B => 4+ // 32B => 6+ //const uint64_t swizzle_enum = 4;+ constexpr uint64_t swizzle_enum =+ (K_BLOCK==256) ? 2 :+ (K_BLOCK==128) ? 4 :+ 6;+auto make_desc_AB = [](uint32_t addr) -> uint64_t {- constexpr int SBO = 8 * 128;- return desc_encode(addr) | (desc_encode(SBO) << 32ULL) | (1ULL << 46ULL) | (2ULL << 61ULL);+ constexpr int SBO = 8 * K_BLOCK/2;+ return desc_encode(addr) | (desc_encode(SBO) << 32ULL) | (1ULL << 46ULL) | (swizzle_enum << 61ULL);};// no swizzlingauto make_desc_SF = [](uint32_t addr) -> uint64_t {⋯ 99 unchanged lines// epilogue and storembarrier_wait(main_loop_mbar_addr, 0);asm volatile("tcgen05.fence::after_thread_sync;");- for (int k=0; k<2; ++k) {- tcgen05_ld_32(d1_regs, ((warpID_4*32) << 16) + k*32);- tcgen05_ld_32(d2_regs, ((warpID_4*32) << 16) + N_BLOCK + k*32);- asm volatile("tcgen05.wait::ld.sync.aligned;");- for (int i=0; i<32; ++i) {- d1_regs[i] = my_sigmoidf(d1_regs[i]) * d1_regs[i];- d1_regs[i] = d1_regs[i] * d2_regs[i];- }+ tcgen05_ld_64(d1_regs, ((warpID_4*32) << 16));+ tcgen05_ld_64(d2_regs, ((warpID_4*32) << 16) + N_BLOCK);+ asm volatile("tcgen05.wait::ld.sync.aligned;");#pragma unroll- for (int i=0; i<16; ++i) {- reinterpret_cast<half2 *>(c_regs)[i] = __float22half2_rn(reinterpret_cast<float2*>(d1_regs)[i]);+ for (int i=0; i<64; ++i) {+ //d1_regs[i] = my_sigmoidf(d1_regs[i]) * d1_regs[i];+ d1_regs[i] = tanh_silu(d1_regs[i], d2_regs[i]);+ }+ /*+ #pragma unroll+ for (int i=0; i<64; ++i) {+ d1_regs[i] *= d2_regs[i];+ }+ */+ #pragma unroll+ for (int i=0; i<32; ++i) {+ reinterpret_cast<half2 *>(c_regs)[i] = __float22half2_rn(reinterpret_cast<float2*>(d1_regs)[i]);+ }+ // store to C+ for (int j=0; j<64; ++j) {+ if ((global_m_id < M) && (global_n_id + j < N)) {+ C[(global_n_id+j)*M + global_m_id] = c_regs[j];+ //C[(global_m_id)*N + global_n_id+j] = c_regs[j];}- // store to C- // each thread has 32 fp16 values, not contiguous sad- uint4 *output_addr = reinterpret_cast<uint4*>(C + global_m_id*N + global_n_id + k*32);- //half *output_addr = C + global_m_id*N + global_n_id + k*32;- for (int j=0; j<4; ++j) {- if (global_m_id < M && (global_n_id+j*8 + k*32 < N)) {- //output_addr[j] = c_regs[j];- output_addr[j] = reinterpret_cast<uint4*>(c_regs)[j];- }- }- __syncwarp();}}__syncthreads();⋯ 68 unchanged linesif (err != cudaSuccess) {throw std::runtime_error(cudaGetErrorString(err));}- return C;+ auto C_transpose = C.view({N, M, 1}).transpose(0, 1);+ //auto C_transpose2 = C.transpose(0, 1);+ /*+ std::cout << "C shape/strides: " << C.sizes() << ", " << C.strides() << std::endl;+ std::cout << "C_transpose shape/strides: " << C_transpose.sizes() << ", " << C_transpose.strides() << std::endl;+ std::cout << "C_transpose2 shape/strides: " << C_transpose2.sizes() << ", " << C_transpose2.strides() << std::endl;++ std::cout << "C[10,15]: " << C.index({10,15,0}) << std::endl;+ std::cout << "C_transpose[10,15]: " << C_transpose.index({10,15,0}) << std::endl;+ */++ return C_transpose;}"""⋯ 91 unchanged lines_, _, l = c.shape#print(sfa.shape, sfa.stride())#print(f"SFA[0,0:32,0]: {sfa[0,:32,0].reshape(-1,2)}")- dual_gemm_cuda(a, b1, b2, sfa_permuted, sfb1_permuted, sfb2_permuted, c)+ out = dual_gemm_cuda(a, b1, b2, sfa_permuted, sfb1_permuted, sfb2_permuted, c)#torch.cuda.synchronize()#print(c_ref)- return c+ return out
scrolls · 188 diff lines total
Best evidence level for this revision: reported
JSON