submission 215834
Quantizr · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 556 lines, June 9 Researcher Reciprocity License v1.0.
yuh.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-nvfp4-dual-gemm-215834?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:361dec90ebe7f5b11cdf0fc6c78d99ea38a67449cd598568f6121ac2a679f877
license declaredunknown
license concludedunknown
authorsQuantizr
imported2026-08-26
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
mbarrier
__device__ __forceinline__ void mbarrier_init(int mbar_addr, int count) {shared-memory
extern __shared__ __align__(1024) char smem_ptr[];split-k
constexpr int SPLIT_K_C = 1;tcgen05
asm volatile("tcgen05.cp.cta_group::1.32x128b.warpx4 [%0], %1;" :: "r"(taddr), "l"(s_desc));tile-k = 256
static_assert(BK == 256, "This hardcoded build uses BK=256 only");tma
"cp.async.bulk.shared::cta.global.mbarrier::complete_tx::bytes.L2::cache_hint "vector-width = half2
half2 h0 = __floats2half2_rn((x0 * s0) * y0, (x1 * s1) * y1);Kernel source
yuh.py556 lines
import torch
from torch.utils.cpp_extension import load_inline
from task import input_t, output_t
CUDA_SRC = r"""
#include <cuda.h>
#include <cuda_runtime.h>
#include <cudaTypedefs.h>
#include <cuda_fp16.h>
#include <torch/library.h>
#include <ATen/core/Tensor.h>
constexpr int WARP_SIZE = 32;
constexpr int MMA_K = 64;
constexpr int SPLIT_K_C = 1;
constexpr uint64_t EVICT_NORMAL = 0x1000000000000000ULL;
constexpr uint64_t EVICT_FIRST = 0x12F0000000000000ULL;
constexpr uint64_t EVICT_LAST = 0x14F0000000000000ULL;
constexpr int SWZ_DESC_128B = 2;
__device__ __forceinline__ uint64_t desc_encode(uint64_t x) { return (x & 0x3'FFFFULL) >> 4ULL; }
__device__ __forceinline__ uint32_t elect_sync() {
uint32_t pred = 0;
asm volatile("{\n\t.reg .pred %px;\n\telect.sync _|%px, %1;\n\t@%px mov.s32 %0, 1;\n\t}"
: "+r"(pred) : "r"(0xFFFFFFFF));
return pred;
}
__device__ __forceinline__ void mbarrier_init(int mbar_addr, int count) {
asm volatile("mbarrier.init.shared::cta.b64 [%0], %1;" :: "r"(mbar_addr), "r"(count));
}
__device__ __forceinline__ void mbarrier_wait(int mbar_addr, int phase) {
uint32_t ticks = 0x989680;
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__ __forceinline__ 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__ __forceinline__ 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__ __forceinline__ void tcgen05_cp_nvfp4(int taddr, uint64_t s_desc) {
asm volatile("tcgen05.cp.cta_group::1.32x128b.warpx4 [%0], %1;" :: "r"(taddr), "l"(s_desc));
}
__device__ __forceinline__ void tcgen05_mma_nvfp4(
int d_tmem,
uint64_t a_desc, uint64_t b_desc, uint32_t i_desc,
int scale_A_tmem, int scale_B_tmem,
int enable_input_d) {
asm volatile("{\n\t.reg .pred p;\n\t"
"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}"
:: "r"(d_tmem), "l"(a_desc), "l"(b_desc), "r"(i_desc),
"r"(scale_A_tmem), "r"(scale_B_tmem), "r"(enable_input_d));
}
// ---- tcgen05 ld for N-major epilogue ----
struct SHAPE { static constexpr char _16x256b[] = ".16x256b"; };
struct NUM { static constexpr char x8[] = ".x8"; static constexpr char x16[] = ".x16"; };
template <const char *SHAPE_S, const char *NUM_S>
__device__ __forceinline__ 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_S), "C"(NUM_S)
);
}
template <const char *SHAPE_S, const char *NUM_S>
__device__ __forceinline__ 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_S), "C"(NUM_S)
);
}
__device__ __forceinline__ void tcgen05_ld_16x256bx8(float *tmp, int row, int col) {
tcgen05_ld_32regs<SHAPE::_16x256b, NUM::x8>(tmp, row, col);
}
__device__ __forceinline__ void tcgen05_ld_16x256bx16(float *tmp, int row, int col) {
tcgen05_ld_64regs<SHAPE::_16x256b, NUM::x16>(tmp, row, col);
}
// ---- TMA map encode ----
static inline void check_cu(CUresult err) {
if (err == CUDA_SUCCESS) return;
const char* msg = "unknown";
cuGetErrorString(err, &msg);
TORCH_CHECK(false, "CUDA Driver error: ", msg);
}
template<int BK>
void init_AB_tmap(CUtensorMap *tmap, const char *ptr, uint64_t g_h, uint64_t g_w, uint32_t s_h, uint32_t s_w) {
static_assert(BK == 256, "This hardcoded build uses BK=256 only");
constexpr int INNER = 256;
constexpr int INNER_BYTES = INNER / 2;
uint64_t gDim[3] = {(uint64_t)INNER, g_h, g_w / (uint64_t)INNER};
uint64_t gStrides[2] = {g_w / 2, (uint64_t)INNER_BYTES};
uint32_t bDim[3] = {(uint32_t)INNER, s_h, (uint32_t)(s_w / (uint32_t)INNER)};
uint32_t eStrides[3] = {1, 1, 1};
check_cu(cuTensorMapEncodeTiled(
tmap,
CUtensorMapDataType::CU_TENSOR_MAP_DATA_TYPE_16U4_ALIGN8B,
3,
(void *)ptr,
gDim,
gStrides,
bDim,
eStrides,
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));
}
template<int BM>
struct CfgT {
static constexpr int EPIL_THREADS = BM;
static constexpr int TB_SIZE = EPIL_THREADS + 2 * WARP_SIZE;
static constexpr int TMA_WARP = EPIL_THREADS / WARP_SIZE;
static constexpr int MMA_WARP = TMA_WARP + 1;
};
template<int K_FIXED, int BM, int BN, int BK, int NUM_STAGES>
__global__ __launch_bounds__(CfgT<BM>::TB_SIZE)
void kernel(
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,
int M, int N
) {
static_assert(BM == 128);
static_assert(BN == 64 || BN == 128);
static_assert(BK == 256);
static_assert(K_FIXED % BK == 0);
const int tid = threadIdx.x;
const int warp_id = tid / WARP_SIZE;
const int lane_id = tid & 31;
constexpr int A_size = BM * BK / 2;
constexpr int B_size = BN * BK / 2;
constexpr int SFA_size = 128 * BK / 16;
constexpr int SFB_size = 128 * BK / 16;
constexpr int STAGE_SIZE = A_size + 2 * B_size + SFA_size + 2 * SFB_size;
extern __shared__ __align__(1024) char smem_ptr[];
const int smem = static_cast<int>(__cvta_generic_to_shared(smem_ptr));
__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;
__shared__ int tmem_base;
constexpr int ACC1_tmem = 0;
constexpr int ACC2_tmem = BN;
constexpr int SFA_tmem = BN * 2;
constexpr int SFB1_tmem = SFA_tmem + 4 * (BK / MMA_K);
constexpr int SFB2_tmem = SFB1_tmem + 4 * (BK / MMA_K);
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;");
}
if (warp_id == 1) {
const int tmem_base_smem = static_cast<int>(__cvta_generic_to_shared(&tmem_base));
asm volatile("tcgen05.alloc.cta_group::1.sync.aligned.shared::cta.b32 [%0], %1;"
:: "r"(tmem_base_smem), "r"(BN * 4));
}
__syncthreads();
constexpr int num_iters = K_FIXED / BK;
constexpr int INNER = 256;
// Scope the whole tile body to shorten live ranges (key “improvement”)
auto run_one_tile = [&](int bid) {
const int grid_n = N / BN;
const int bid_m = bid / grid_n;
const int bid_n = bid % grid_n;
const int off_m = bid_m * BM;
const int off_n = bid_n * BN;
// Producer warp
if (warp_id == CfgT<BM>::TMA_WARP && elect_sync()) {
uint64_t cache_A = (M > N) ? EVICT_FIRST : EVICT_LAST;
uint64_t cache_B = (M > N) ? EVICT_LAST : 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 B1_smem = A_smem + A_size;
const int B2_smem = B1_smem + B_size;
const int SFA_smem = B2_smem + B_size;
const int SFB1_smem = SFA_smem + SFA_size;
const int SFB2_smem = SFB1_smem + SFB_size;
const int off_k = iter_k * BK;
tma_3d_gmem2smem(A_smem, &A_tmap, 0, off_m, off_k / INNER, mbar_addr, cache_A);
tma_3d_gmem2smem(B1_smem, &B1_tmap, 0, off_n, off_k / INNER, mbar_addr, cache_B);
tma_3d_gmem2smem(B2_smem, &B2_tmap, 0, off_n, off_k / INNER, mbar_addr, cache_B);
const int rest_k = K_FIXED / 64;
const char *SFA_src = SFA_ptr + ((off_m / 128) * rest_k + off_k / 64) * 512;
const char *SFB1_src = SFB1_ptr + ((off_n / 128) * rest_k + off_k / 64) * 512;
const char *SFB2_src = SFB2_ptr + ((off_n / 128) * rest_k + off_k / 64) * 512;
tma_gmem2smem(SFA_smem, SFA_src, SFA_size, mbar_addr, cache_A);
tma_gmem2smem(SFB1_smem, SFB1_src, SFB_size, mbar_addr, cache_B);
tma_gmem2smem(SFB2_smem, SFB2_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");
};
#pragma unroll
for (int i = 0; i < NUM_STAGES; i++) issue_tma(i, i);
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);
}
}
// Consumer warp
else if (warp_id == CfgT<BM>::MMA_WARP && elect_sync()) {
constexpr int MMA_N = BN;
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);
const int taddr = tmem_base;
constexpr int SWZ_MODE = SWZ_DESC_128B;
constexpr int SWZ_BYTES = 128;
auto make_desc_AB = [](int addr) -> uint64_t {
const int SBO = 8 * SWZ_BYTES;
return desc_encode((uint64_t)addr) |
(desc_encode((uint64_t)SBO) << 32ULL) |
(1ULL << 46ULL) |
((uint64_t)SWZ_MODE << 61ULL);
};
auto make_desc_SF = [](int addr) -> uint64_t {
const int SBO = 8 * 16;
return desc_encode((uint64_t)addr) |
(desc_encode((uint64_t)SBO) << 32ULL) |
(1ULL << 46ULL);
};
constexpr uint64_t AB_step = (32ULL >> 4ULL);
constexpr uint64_t SF_step = (512ULL >> 4ULL);
const int scale_A_blk = (bid_m % (128 / BM)) * (BM / 32);
const int scale_B_blk = (bid_n % (128 / BN)) * (BN / 32);
constexpr int SEG_K = 256;
constexpr int SEG_MMA_STEPS = SEG_K / MMA_K; // 4
constexpr int NUM_SEGS = BK / SEG_K; // 1
constexpr int A_SEG_BYTES = BM * SEG_K / 2;
constexpr int B_SEG_BYTES = BN * SEG_K / 2;
constexpr int SF_SEG_BYTES = 128 * SEG_K / 16;
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 B1_smem = A_smem + A_size;
const int B2_smem = B1_smem + B_size;
const int SFA_smem = B2_smem + B_size;
const int SFB1_smem = SFA_smem + SFA_size;
const int SFB2_smem = SFB1_smem + SFB_size;
#pragma unroll
for (int seg = 0; seg < NUM_SEGS; seg++) {
const int A_seg_smem = A_smem + seg * A_SEG_BYTES;
const int B1_seg_smem = B1_smem + seg * B_SEG_BYTES;
const int B2_seg_smem = B2_smem + seg * B_SEG_BYTES;
const int SFA_seg_smem = SFA_smem + seg * SF_SEG_BYTES;
const int SFB1_seg_smem = SFB1_smem + seg * SF_SEG_BYTES;
const int SFB2_seg_smem = SFB2_smem + seg * SF_SEG_BYTES;
const uint64_t SF_desc0 = make_desc_SF(0);
const uint64_t SFA_desc_base = SF_desc0 + ((uint64_t)SFA_seg_smem >> 4ULL);
const uint64_t SFB1_desc_base = SF_desc0 + ((uint64_t)SFB1_seg_smem >> 4ULL);
const uint64_t SFB2_desc_base = SF_desc0 + ((uint64_t)SFB2_seg_smem >> 4ULL);
const uint64_t A_desc0 = make_desc_AB(A_seg_smem);
const uint64_t B1_desc0 = make_desc_AB(B1_seg_smem);
const uint64_t B2_desc0 = make_desc_AB(B2_seg_smem);
#pragma unroll
for (int kk = 0; kk < SEG_MMA_STEPS; kk++) {
const int k_global = seg * SEG_MMA_STEPS + kk;
tcgen05_cp_nvfp4(taddr + SFA_tmem + k_global * 4, SFA_desc_base + (uint64_t)kk * SF_step);
tcgen05_cp_nvfp4(taddr + SFB1_tmem + k_global * 4, SFB1_desc_base + (uint64_t)kk * SF_step);
tcgen05_cp_nvfp4(taddr + SFB2_tmem + k_global * 4, SFB2_desc_base + (uint64_t)kk * SF_step);
}
#pragma unroll
for (int kk = 0; kk < SEG_MMA_STEPS; kk++) {
const int k_global = seg * SEG_MMA_STEPS + kk;
const uint64_t a_desc = A_desc0 + (uint64_t)kk * AB_step;
const uint64_t b1_desc = B1_desc0 + (uint64_t)kk * AB_step;
const uint64_t b2_desc = B2_desc0 + (uint64_t)kk * AB_step;
const int scale_A_tmem = (taddr + SFA_tmem) + k_global * 4 + scale_A_blk;
const int scale_B1_tmem = (taddr + SFB1_tmem) + k_global * 4 + scale_B_blk;
const int scale_B2_tmem = (taddr + SFB2_tmem) + k_global * 4 + scale_B_blk;
const int enable_input_d = (k_global == 0) ? iter_k : 1;
tcgen05_mma_nvfp4(taddr + ACC1_tmem, a_desc, b1_desc, i_desc, scale_A_tmem, scale_B1_tmem, enable_input_d);
tcgen05_mma_nvfp4(taddr + ACC2_tmem, a_desc, b2_desc, i_desc, scale_A_tmem, scale_B2_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");
}
// Epilogue
if (tid < BM) {
mbarrier_wait(mainloop_mbar_addr, 0);
asm volatile("tcgen05.fence::after_thread_sync;");
#pragma unroll
for (int m16 = 0; m16 < 32; m16 += 16) {
float a_tmp[BN / 2];
float b_tmp[BN / 2];
const int tmem_row = warp_id * 32 + m16;
if constexpr (BN == 64) {
tcgen05_ld_16x256bx8(a_tmp, tmem_row, ACC1_tmem);
tcgen05_ld_16x256bx8(b_tmp, tmem_row, ACC2_tmem);
} else {
tcgen05_ld_16x256bx16(a_tmp, tmem_row, ACC1_tmem);
tcgen05_ld_16x256bx16(b_tmp, tmem_row, ACC2_tmem);
}
asm volatile("tcgen05.wait::ld.sync.aligned;");
const int row0 = off_m + warp_id * 32 + m16 + (lane_id / 4);
const int row1 = row0 + 8;
const int lane2 = (lane_id % 4) * 2;
#pragma unroll
for (int i = 0; i < BN / 8; i++) {
const int col = off_n + i * 8 + lane2;
const int j = i * 4;
float x0 = a_tmp[j + 0], x1 = a_tmp[j + 1];
float x2 = a_tmp[j + 2], x3 = a_tmp[j + 3];
float y0 = b_tmp[j + 0], y1 = b_tmp[j + 1];
float y2 = b_tmp[j + 2], y3 = b_tmp[j + 3];
float s0 = 1.0f / (1.0f + __expf(-x0));
float s1 = 1.0f / (1.0f + __expf(-x1));
float s2 = 1.0f / (1.0f + __expf(-x2));
float s3 = 1.0f / (1.0f + __expf(-x3));
half2 h0 = __floats2half2_rn((x0 * s0) * y0, (x1 * s1) * y1);
half2 h1 = __floats2half2_rn((x2 * s2) * y2, (x3 * s3) * y3);
*reinterpret_cast<half2*>(C_ptr + row0 * N + col) = h0;
*reinterpret_cast<half2*>(C_ptr + row1 * N + col) = h1;
}
}
asm volatile("bar.sync 1, %0;" :: "r"(BM) : "memory");
if (tid == 0) {
const int taddr = tmem_base;
asm volatile("tcgen05.dealloc.cta_group::1.sync.aligned.b32 %0, %1;"
:: "r"(taddr), "r"(BN * 4));
}
}
};
// Non-persistent path (still benefits from scoped tile body)
run_one_tile((int)blockIdx.y);
}
template<int K_FIXED, int BM, int BN, int BK, int NUM_STAGES>
void launch_one(
const at::Tensor& A,
const at::Tensor& B1,
const at::Tensor& B2,
const at::Tensor& SFAp,
const at::Tensor& SFB1p,
const at::Tensor& SFB2p,
at::Tensor& C
) {
const int M = (int)A.size(0);
const int N = (int)B1.size(0);
CUtensorMap A_tmap, B1_tmap, B2_tmap;
init_AB_tmap<BK>(&A_tmap, (const char*)A.data_ptr(), (uint64_t)M, (uint64_t)K_FIXED, (uint32_t)BM, (uint32_t)BK);
init_AB_tmap<BK>(&B1_tmap, (const char*)B1.data_ptr(), (uint64_t)N, (uint64_t)K_FIXED, (uint32_t)BN, (uint32_t)BK);
init_AB_tmap<BK>(&B2_tmap, (const char*)B2.data_ptr(), (uint64_t)N, (uint64_t)K_FIXED, (uint32_t)BN, (uint32_t)BK);
const int tiles = (M / BM) * (N / BN);
dim3 grid(1, tiles);
int A_sz = (BM * BK / 2);
int B_sz = (BN * BK / 2);
int SFA_sz = (128 * (BK / 16));
int SFB_sz = (128 * (BK / 16));
int stage_sz = A_sz + 2 * B_sz + SFA_sz + 2 * SFB_sz;
int smem_size = stage_sz * NUM_STAGES;
auto k = kernel<K_FIXED, BM, BN, BK, NUM_STAGES>;
if (smem_size > 48000) cudaFuncSetAttribute(k, cudaFuncAttributeMaxDynamicSharedMemorySize, smem_size);
constexpr int TB = CfgT<BM>::TB_SIZE;
k<<<grid, TB, smem_size>>>(
A_tmap, B1_tmap, B2_tmap,
(const char*)SFAp.data_ptr(),
(const char*)SFB1p.data_ptr(),
(const char*)SFB2p.data_ptr(),
(half*)C.data_ptr(),
M, N);
}
// Hardcoded dispatch in C++ (cfg15/cfg19 equivalents)
at::Tensor run(
const at::Tensor& A, const at::Tensor& B1, const at::Tensor& B2,
const at::Tensor& SFAp, const at::Tensor& SFB1p, const at::Tensor& SFB2p,
at::Tensor& C
) {
const int M = (int)C.size(0);
const int N = (int)C.size(1);
const int K = (int)(A.size(1) * 2);
// Best from your sweep:
// cfg15: BM128 BN64 BK256 NS5
// cfg19: BM128 BN128 BK256 NS4
if (K == 7168 && M == 256 && N == 4096) { launch_one<7168, 128, 64, 256, 5>(A,B1,B2,SFAp,SFB1p,SFB2p,C); return C; }
if (K == 4096 && M == 256 && N == 3072) { launch_one<4096, 128, 64, 256, 5>(A,B1,B2,SFAp,SFB1p,SFB2p,C); return C; }
if (K == 7168 && M == 512 && N == 4096) { launch_one<7168, 128, 128, 256, 4>(A,B1,B2,SFAp,SFB1p,SFB2p,C); return C; }
if (K == 7168 && M == 512 && N == 3072) { launch_one<7168, 128, 128, 256, 4>(A,B1,B2,SFAp,SFB1p,SFB2p,C); return C; }
if (K == 7168) { launch_one<7168, 128, 64, 256, 5>(A,B1,B2,SFAp,SFB1p,SFB2p,C); return C; }
if (K == 2304) { launch_one<2304, 128, 64, 256, 5>(A,B1,B2,SFAp,SFB1p,SFB2p,C); return C; }
if (K == 2048) { launch_one<2048, 128, 64, 256, 5>(A,B1,B2,SFAp,SFB1p,SFB2p,C); return C; }
if (K == 1536) { launch_one<1536, 128, 64, 256, 5>(A,B1,B2,SFAp,SFB1p,SFB2p,C); return C; }
if (K == 512) { launch_one<512, 128, 64, 256, 5>(A,B1,B2,SFAp,SFB1p,SFB2p,C); return C; }
if (K == 256) { launch_one<256, 128, 64, 256, 5>(A,B1,B2,SFAp,SFB1p,SFB2p,C); return C; }
TORCH_CHECK(false, "Unsupported shape M=", M, " N=", N, " K=", K);
}
TORCH_LIBRARY(dual_fixed_scoped, m) {
m.def("run(Tensor A, Tensor B1, Tensor B2, Tensor SFAp, Tensor SFB1p, Tensor SFB2p, Tensor(a!) C) -> Tensor");
m.impl("run", &run);
}
"""
load_inline(
name="dual_fixed_scoped",
cpp_sources="",
cuda_sources=CUDA_SRC,
verbose=True,
is_python_module=False,
no_implicit_headers=True,
extra_cuda_cflags=[
"-O3",
"-gencode=arch=compute_100a,code=sm_100a",
"--use_fast_math",
"--expt-relaxed-constexpr",
"--relocatable-device-code=false",
"-lineinfo",
"-Xptxas=-v",
"-std=c++17",
],
extra_ldflags=["-lcuda"],
)
_run = torch.ops.dual_fixed_scoped.run
def custom_kernel(data: input_t) -> output_t:
a, b1, b2, _, _, _, sfa_perm, sfb1_perm, sfb2_perm, c = data
_run(a, b1, b2, sfa_perm, sfb1_perm, sfb2_perm, c)
return c
scrolls · 556 lines total
Source code from GPU Mode and the KernelBot dataset · June 9 Researcher Reciprocity License v1.0
Best evidence level for this revision: reported
JSON