"""FP8 e4m3 GEMM for B200 (SM100): y = (x @ w.T) * weight_scale, bf16 out. Hand-written CUDA C++ tcgen05 (5th-gen tensor core) kernel: persistent warp-specialized pipeline (TMA loader warp + MMA umpire warp + epilogue warpgroup), fp8 x fp8 MMA with fp32 accumulation in tensor memory, per-channel dequant scale in an SMEM-staged epilogue. No vendor GEMM calls anywhere. """ import torch import torch.nn as nn from torch.utils.cpp_extension import load_inline OP_TYPE = "gemm" SUPPORTED_PRECISIONS = ["fp8_e4m3"] HARDWARE_REQUIRED = ["B200"] E4M3_MAX = 448.0 _CUDA_SRC = r""" #include #include #include #include #include #include #define DEV __device__ __forceinline__ DEV static uint32_t elect_one_sync() { uint32_t pred = 0; asm volatile( "{\n" ".reg .pred P1;\n" "elect.sync _|P1, %1;\n" "selp.b32 %0, 1, 0, P1;\n" "}\n" : "=r"(pred) : "r"(0xffffffff)); return pred; } DEV static void mbar_init(uint64_t* bar, uint32_t count) { asm volatile("mbarrier.init.shared::cta.b64 [%0], %1;\n" : : "r"((uint32_t)__cvta_generic_to_shared(bar)), "r"(count)); } DEV static void mbar_expect(uint64_t* bar, uint32_t bytes) { asm volatile("mbarrier.arrive.expect_tx.shared::cta.b64 _, [%0], %1;\n" : : "r"((uint32_t)__cvta_generic_to_shared(bar)), "r"(bytes)); } DEV static void mbar_arrive(uint64_t* bar) { asm volatile("mbarrier.arrive.shared::cta.b64 _, [%0];\n" : : "r"((uint32_t)__cvta_generic_to_shared(bar))); } DEV static void mbar_wait(uint64_t* bar, uint32_t phase) { asm volatile( "{\n" ".reg .pred P1;\n" "WAITLOOP:\n" "mbarrier.try_wait.parity.shared::cta.b64 P1, [%0], %1;\n" "@P1 bra DONE;\n" "bra WAITLOOP;\n" "DONE:\n" "}\n" : : "r"((uint32_t)__cvta_generic_to_shared(bar)), "r"(phase)); } DEV static void tma_load_2d(const void* smem_dst, const CUtensorMap* map, int32_t c0, int32_t c1, uint64_t* bar) { asm volatile( "cp.async.bulk.tensor.2d.shared::cluster.global.mbarrier::complete_tx::bytes" " [%0], [%1, {%3, %4}], [%2];\n" : : "r"((uint32_t)__cvta_generic_to_shared(smem_dst)), "l"((uint64_t)map), "r"((uint32_t)__cvta_generic_to_shared(bar)), "r"(c0), "r"(c1) : "memory"); } DEV static void tma_store_2d(const CUtensorMap* map, const void* smem_src, int32_t c0, int32_t c1) { asm volatile( "cp.async.bulk.tensor.2d.global.shared::cta.bulk_group [%0, {%2, %3}], [%1];\n" : : "l"((uint64_t)map), "r"((uint32_t)__cvta_generic_to_shared(smem_src)), "r"(c0), "r"(c1) : "memory"); } DEV static void tma_store_commit() { asm volatile("cp.async.bulk.commit_group;"); } template DEV static void tma_store_wait() { asm volatile("cp.async.bulk.wait_group.read %0;" ::"n"(N)); } DEV static void fence_async_shared() { asm volatile("fence.proxy.async.shared::cta;"); } // smem matrix descriptor (UMMA). layout_type: 2=SW128, 4=SW64, 6=SW32. DEV static uint64_t smem_desc(uint32_t saddr, uint32_t sbo, uint32_t swz_type) { uint64_t d = ((uint64_t)(saddr >> 4) & 0x3FFFULL); d |= ((uint64_t)(sbo >> 4)) << 32; d |= 1ULL << 46; // version for blackwell d |= (uint64_t)swz_type << 61; return d; } DEV static void umma_f8(uint64_t da, uint64_t db, uint32_t tmem_d, uint32_t idesc, uint32_t accum) { asm volatile( "{\n\t" ".reg .pred p;\n\t" "setp.ne.b32 p, %4, 0;\n\t" "tcgen05.mma.cta_group::1.kind::f8f6f4 [%0], %1, %2, %3, {%5, %6, %7, %8}, p; \n\t" "}\n" : : "r"(tmem_d), "l"(da), "l"(db), "r"(idesc), "r"(accum), "r"(0), "r"(0), "r"(0), "r"(0)); } DEV static void umma_commit(uint64_t* bar) { asm volatile( "tcgen05.commit.cta_group::1.mbarrier::arrive::one.shared::cluster.b64 [%0];\n" : : "r"((uint32_t)__cvta_generic_to_shared(bar))); } DEV static void tmem_alloc(uint32_t* smem_out, uint32_t ncols) { asm volatile("tcgen05.alloc.cta_group::1.sync.aligned.shared::cta.b32 [%0], %1;\n" : : "r"((uint32_t)__cvta_generic_to_shared(smem_out)), "r"(ncols)); } DEV static void tmem_dealloc(uint32_t taddr, uint32_t ncols) { asm volatile("tcgen05.dealloc.cta_group::1.sync.aligned.b32 %0, %1;\n" ::"r"(taddr), "r"(ncols)); asm volatile("tcgen05.relinquish_alloc_permit.cta_group::1.sync.aligned;"); } DEV static void tmem_ld64(float* vals, uint32_t taddr) { uint32_t* v = (uint32_t*)vals; asm volatile("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" : "=r"(v[0]), "=r"(v[1]), "=r"(v[2]), "=r"(v[3]), "=r"(v[4]), "=r"(v[5]), "=r"(v[6]), "=r"(v[7]), "=r"(v[8]), "=r"(v[9]), "=r"(v[10]), "=r"(v[11]), "=r"(v[12]), "=r"(v[13]), "=r"(v[14]), "=r"(v[15]), "=r"(v[16]), "=r"(v[17]), "=r"(v[18]), "=r"(v[19]), "=r"(v[20]), "=r"(v[21]), "=r"(v[22]), "=r"(v[23]), "=r"(v[24]), "=r"(v[25]), "=r"(v[26]), "=r"(v[27]), "=r"(v[28]), "=r"(v[29]), "=r"(v[30]), "=r"(v[31]), "=r"(v[32]), "=r"(v[33]), "=r"(v[34]), "=r"(v[35]), "=r"(v[36]), "=r"(v[37]), "=r"(v[38]), "=r"(v[39]), "=r"(v[40]), "=r"(v[41]), "=r"(v[42]), "=r"(v[43]), "=r"(v[44]), "=r"(v[45]), "=r"(v[46]), "=r"(v[47]), "=r"(v[48]), "=r"(v[49]), "=r"(v[50]), "=r"(v[51]), "=r"(v[52]), "=r"(v[53]), "=r"(v[54]), "=r"(v[55]), "=r"(v[56]), "=r"(v[57]), "=r"(v[58]), "=r"(v[59]), "=r"(v[60]), "=r"(v[61]), "=r"(v[62]), "=r"(v[63]) : "r"(taddr)); asm volatile("tcgen05.wait::ld.sync.aligned;"); } DEV static constexpr uint32_t make_idesc(int M, int N) { return (1u << 4) | ((uint32_t)(N >> 3) << 17) | ((uint32_t)(M >> 4) << 24); } // ------------------------------------------------------------------ scheduling struct TileScheduler { int start_pid, num_pid_m, num_pid_n, num_pid_in_group, num_pid, num_progs; int group_m; __device__ void init(int M, int N, int BM, int BN, int gm) { start_pid = blockIdx.x; num_progs = gridDim.x; num_pid_m = (M + BM - 1) / BM; num_pid_n = (N + BN - 1) / BN; num_pid_in_group = gm * num_pid_n; num_pid = num_pid_m * num_pid_n; group_m = gm; } __device__ int num_tiles() const { return (num_pid - start_pid + num_progs - 1) / num_progs; } __device__ void get_tile(int idx, int& pm, int& pn) const { int tile_id = start_pid + idx * num_progs; int group_id = tile_id / num_pid_in_group; int first_pid_m = group_id * group_m; int group_size = min(num_pid_m - first_pid_m, group_m); pm = first_pid_m + (tile_id % group_size); pn = (tile_id % num_pid_in_group) / group_size; } }; // ------------------------------------------------------------------ kernel template __global__ void __launch_bounds__(192, 1) k_gemm(const __grid_constant__ CUtensorMap a_map, const __grid_constant__ CUtensorMap b_map, const __grid_constant__ CUtensorMap c_map, const float* __restrict__ scale, int M, int N, int K) { constexpr int A_BYTES = BM * BK; constexpr int B_BYTES = BN * BK; constexpr int C_COLS = 64; // bf16 columns per c box constexpr int C_ELEMS = BM * C_COLS; // bf16 elements per c buffer constexpr int N_CHUNKS = BN / C_COLS; constexpr int EPI_WARPS = 4; constexpr int EPI_THREADS = EPI_WARPS * 32; // m64 accumulator packs rows into 16-lane groups; the two acc buffers are // interleaved in lane halves: buf b at lane offset 16*b (same columns). // m128 accumulator: buf b at column offset b * BN. constexpr int SWZ_TYPE = (BK == 128) ? 2 : 4; // 128B or 64B swizzle extern __shared__ __align__(1024) uint8_t smem[]; uint8_t* a_smem = smem; // STAGES * A_BYTES uint8_t* b_smem = a_smem + STAGES * A_BYTES; // STAGES * B_BYTES __nv_bfloat16* c_smem = (__nv_bfloat16*)(b_smem + STAGES * B_BYTES); // 2 * C_ELEMS float* s_smem = (float*)(c_smem + 2 * C_ELEMS); // BN floats __shared__ __align__(8) uint64_t full_bars[STAGES], empty_bars[STAGES]; __shared__ __align__(8) uint64_t acc_full[2], acc_empty[2]; __shared__ uint32_t tmem_addr_sh; const int warp = threadIdx.x >> 5; const int lane = threadIdx.x & 31; if (warp == EPI_WARPS) { tmem_alloc(&tmem_addr_sh, 512); } else if (warp == EPI_WARPS + 1 && lane == 0) { for (int i = 0; i < STAGES; ++i) { mbar_init(&full_bars[i], 1); mbar_init(&empty_bars[i], 1); } mbar_init(&acc_full[0], 1); mbar_init(&acc_full[1], 1); mbar_init(&acc_empty[0], EPI_WARPS); mbar_init(&acc_empty[1], EPI_WARPS); } asm volatile("barrier.sync 1, 192;"); const uint32_t taddr_base = tmem_addr_sh; const uint32_t idesc = make_idesc(BM, BN); constexpr uint32_t ACC1_OFF = (BM == 64) ? (16u << 16) : (uint32_t)BN; TileScheduler sched; sched.init(M, N, BM, BN, GM); const int k_iters = (K + BK - 1) / BK; if (warp == EPI_WARPS) { // ---------------------------------------------------------- loader if (elect_one_sync()) { uint32_t e_phase = 1; // first STAGES waits bypass int stage = 0; for (int t = 0; t < sched.num_tiles(); ++t) { int pm, pn; sched.get_tile(t, pm, pn); const int off_m = pm * BM, off_n = pn * BN; for (int k = 0; k < k_iters; ++k) { const int s = stage; mbar_wait(&empty_bars[s], e_phase); mbar_expect(&full_bars[s], A_BYTES + B_BYTES); tma_load_2d(a_smem + s * A_BYTES, &a_map, k * BK, off_m, &full_bars[s]); tma_load_2d(b_smem + s * B_BYTES, &b_map, k * BK, off_n, &full_bars[s]); if (++stage == STAGES) { stage = 0; e_phase ^= 1; } } } } } else if (warp == EPI_WARPS + 1) { // ---------------------------------------------------------- mma umpire if (elect_one_sync()) { uint32_t f_phase = 0, ae_phase = 1; int stage = 0, abuf = 0; for (int t = 0; t < sched.num_tiles(); ++t) { mbar_wait(&acc_empty[abuf], ae_phase); const uint32_t taddr = taddr_base + abuf * ACC1_OFF; for (int k = 0; k < k_iters; ++k) { const int s = stage; mbar_wait(&full_bars[s], f_phase); uint64_t da = smem_desc((uint32_t)__cvta_generic_to_shared(a_smem + s * A_BYTES), 8 * BK, SWZ_TYPE); uint64_t db = smem_desc((uint32_t)__cvta_generic_to_shared(b_smem + s * B_BYTES), 8 * BK, SWZ_TYPE); for (int kk = 0; kk < BK / 32; ++kk) { umma_f8(da + kk * 2, db + kk * 2, taddr, idesc, (k > 0 || kk > 0) ? 1u : 0u); } umma_commit(&empty_bars[s]); if (++stage == STAGES) { stage = 0; f_phase ^= 1; } } umma_commit(&acc_full[abuf]); if (++abuf == 2) { abuf = 0; ae_phase ^= 1; } } } } else if (warp < EPI_WARPS) { // ---------------------------------------------------------- epilogue // m128: each warp owns 32 rows. m64: lanes hold 16 rows per warp in the // 16-lane half selected by the accumulator buffer index. int row = 0; bool row_ok = true; if (BM == 128) { row = warp * 32 + lane; } else { row_ok = (lane / 16) == 0; // placeholder; abuf-dependent, set in loop } uint32_t af_phase = 0; int abuf = 0; for (int t = 0; t < sched.num_tiles(); ++t) { int pm, pn; sched.get_tile(t, pm, pn); const int off_m = pm * BM, off_n = pn * BN; // preload scales for this tile (BN floats) into smem for (int j = threadIdx.x; j < BN; j += EPI_THREADS) { const int n = off_n + j; s_smem[j] = (n < N) ? scale[n] : 0.f; } mbar_wait(&acc_full[abuf], af_phase); const uint32_t taddr = taddr_base + abuf * ACC1_OFF; for (int c = 0; c < N_CHUNKS; ++c) { float vals[C_COLS]; tmem_ld64(vals, taddr + c * C_COLS); if (c == N_CHUNKS - 1 && lane == 0) { // all tmem reads of this accumulator are done mbar_arrive(&acc_empty[abuf]); } __nv_bfloat16 outv[C_COLS]; const float* sv = s_smem + c * C_COLS; for (int j = 0; j < C_COLS; ++j) { outv[j] = __float2bfloat16(vals[j] * sv[j]); } // wait for the tma store that used this buffer 2 chunks ago if (c >= 2 && warp == 0 && elect_one_sync()) { tma_store_wait<1>(); } asm volatile("barrier.sync 2, %0;" ::"n"(EPI_THREADS)); __nv_bfloat16* dst = c_smem + (c & 1) * C_ELEMS; if (BM == 128) { uint4* d4 = (uint4*)(dst + row * C_COLS); const uint4* s4 = (const uint4*)outv; for (int j = 0; j < C_COLS * 2 / 16; ++j) d4[j] = s4[j]; } else { // m64: lanes 16*abuf..16*abuf+16 of this warp hold rows // warp*16 + (lane - 16*abuf). const int l16 = lane - 16 * abuf; if (l16 >= 0 && l16 < 16) { uint4* d4 = (uint4*)(dst + (warp * 16 + l16) * C_COLS); const uint4* s4 = (const uint4*)outv; for (int j = 0; j < C_COLS * 2 / 16; ++j) d4[j] = s4[j]; } } asm volatile("barrier.sync 2, %0;" ::"n"(EPI_THREADS)); if (warp == 0 && elect_one_sync()) { fence_async_shared(); tma_store_2d(&c_map, c_smem + (c & 1) * C_ELEMS, off_n + c * C_COLS, off_m); tma_store_commit(); } } if (++abuf == 2) { abuf = 0; af_phase ^= 1; } } if (warp == 0 && elect_one_sync()) { tma_store_wait<0>(); } } asm volatile("barrier.sync 3, 192;"); if (warp == EPI_WARPS + 1) { tmem_dealloc(taddr_base, 512); } } // ------------------------------------------------------------------ host __global__ void k_pad2_funnel(const uint32_t* __restrict__ src_a, uint32_t* __restrict__ dst_a, const uint32_t* __restrict__ src_b, uint32_t* __restrict__ dst_b, int K, int dst_words, int src_words_a, int src_words_b, int rows_b) { const int row = blockIdx.x; const bool is_a = row < rows_b; const uint32_t* src = is_a ? src_a : src_b; uint32_t* dst = is_a ? dst_a : dst_b; const int src_words = is_a ? src_words_a : src_words_b; const long long sr = (long long)(is_a ? row : row - rows_b) * K; const int w0 = (int)(sr >> 2); const int sh = (int)(sr & 3) * 8; for (int j = threadIdx.x; j < dst_words; j += blockDim.x) { const int sw = w0 + j; const unsigned a = (sw < src_words) ? src[sw] : 0u; const unsigned b = (sw + 1 < src_words) ? src[sw + 1] : 0u; unsigned w = (a >> sh); if (sh != 0) w |= (b << (32 - sh)); const int byte0 = j * 4; int vb = K - byte0; vb = vb < 0 ? 0 : (vb > 4 ? 4 : vb); const unsigned keep = (vb >= 4) ? 0xFFFFFFFFu : ((1u << (vb * 8)) - 1); dst[(size_t)(is_a ? row : row - rows_b) * dst_words + j] = w & keep; } } // pad kernel: unpack odd-K fp8 rows into u32-aligned rows via funnel shifts // (128-bit stores + shared adjacent-word loads) __global__ void k_pad_funnel(const uint32_t* __restrict__ src, uint4* __restrict__ dst, int K, int dst_words, int src_words) { const int row = blockIdx.x; const long long sr = (long long)row * K; const int w0 = (int)(sr >> 2); const int sh = (int)(sr & 3) * 8; const int n_v = dst_words >> 2; for (int j = threadIdx.x; j < n_v; j += blockDim.x) { const int sw = w0 + (j << 2); const unsigned s0 = (sw < src_words) ? src[sw] : 0u; const unsigned s1 = (sw + 1 < src_words) ? src[sw + 1] : 0u; const unsigned s2 = (sw + 2 < src_words) ? src[sw + 2] : 0u; const unsigned s3 = (sw + 3 < src_words) ? src[sw + 3] : 0u; const unsigned s4 = (sw + 4 < src_words) ? src[sw + 4] : 0u; uint4 v; if (sh == 0) { v.x = s0; v.y = s1; v.z = s2; v.w = s3; } else { const int shr = 32 - sh; v.x = (s0 >> sh) | (s1 << shr); v.y = (s1 >> sh) | (s2 << shr); v.z = (s2 >> sh) | (s3 << shr); v.w = (s3 >> sh) | (s4 << shr); } const int base = (j << 2) * 4; int vb; vb = K - base; vb = vb < 0 ? 0 : (vb > 4 ? 4 : vb); v.x &= (vb >= 4) ? 0xFFFFFFFFu : ((1u << (vb * 8)) - 1); vb = K - (base + 4); vb = vb < 0 ? 0 : (vb > 4 ? 4 : vb); v.y &= (vb >= 4) ? 0xFFFFFFFFu : ((1u << (vb * 8)) - 1); vb = K - (base + 8); vb = vb < 0 ? 0 : (vb > 4 ? 4 : vb); v.z &= (vb >= 4) ? 0xFFFFFFFFu : ((1u << (vb * 8)) - 1); vb = K - (base + 12); vb = vb < 0 ? 0 : (vb > 4 ? 4 : vb); v.w &= (vb >= 4) ? 0xFFFFFFFFu : ((1u << (vb * 8)) - 1); dst[(size_t)row * n_v + j] = v; } } namespace { struct MapCache { std::unordered_map maps; CUtensorMap get(const void* base, uint64_t inner, uint64_t outer, uint32_t box_inner, uint32_t box_outer, CUtensorMapSwizzle swz, CUtensorMapDataType dt, int elembytes) { uint64_t key = ((uint64_t)base) ^ (inner << 1) ^ (outer << 13) ^ ((uint64_t)box_inner << 40) ^ ((uint64_t)box_outer << 48) ^ ((uint64_t)swz << 56) ^ ((uint64_t)dt << 61); auto it = maps.find(key); if (it != maps.end()) return it->second; CUtensorMap map{}; uint64_t dims[2] = {inner, outer}; uint64_t strides[1] = {inner * (uint64_t)elembytes}; uint32_t box[2] = {box_inner, box_outer}; uint32_t estrides[2] = {1, 1}; cuTensorMapEncodeTiled(&map, dt, 2, const_cast(base), dims, strides, box, estrides, CU_TENSOR_MAP_INTERLEAVE_NONE, swz, CU_TENSOR_MAP_L2_PROMOTION_L2_128B, CU_TENSOR_MAP_FLOAT_OOB_FILL_NONE); return map; } }; constexpr CUtensorMapSwizzle CuSw(int bk) { return bk == 128 ? CU_TENSOR_MAP_SWIZZLE_128B : CU_TENSOR_MAP_SWIZZLE_64B; } MapCache g_maps; int g_num_sms = 0; } // namespace template void launch_cfg(const CUtensorMap* a_map, const CUtensorMap* b_map, const CUtensorMap* c_map, const float* scale, int M, int N, int K, cudaStream_t stream) { constexpr int SMEM = STAGES * (BM * BK + BN * BK) + 2 * (BM * 64 * 2) + BN * 4 + 1024; auto kern = k_gemm; static bool attr_set = false; if (!attr_set) { cudaFuncSetAttribute(kern, cudaFuncAttributeMaxDynamicSharedMemorySize, SMEM); attr_set = true; } int tiles_m = (M + BM - 1) / BM, tiles_n = (N + BN - 1) / BN; int grid = std::min(g_num_sms, tiles_m * tiles_n); kern<<>>(*a_map, *b_map, *c_map, scale, M, N, K); } void fp8_run(int64_t x_ptr, int64_t w_ptr, int64_t s_ptr, int64_t y_ptr, int64_t M, int64_t N, int64_t K, int64_t cfg, int64_t stream_ptr) { void* x = (void*)x_ptr; void* w = (void*)w_ptr; if (!g_num_sms) { cudaDeviceProp prop; cudaGetDeviceProperties(&prop, 0); g_num_sms = prop.multiProcessorCount; } cudaStream_t stream = (cudaStream_t)stream_ptr; if (cfg == 0) { CUtensorMap a_map = g_maps.get(x, K, M, 128, 64, CuSw(128), CU_TENSOR_MAP_DATA_TYPE_UINT8, 1); CUtensorMap b_map = g_maps.get(w, K, N, 128, 64, CuSw(128), CU_TENSOR_MAP_DATA_TYPE_UINT8, 1); CUtensorMap c_map = g_maps.get((void*)y_ptr, N, M, 64, 64, CU_TENSOR_MAP_SWIZZLE_NONE, CU_TENSOR_MAP_DATA_TYPE_UINT16, 2); launch_cfg<64, 64, 128, 12, 1>(&a_map, &b_map, &c_map, (const float*)s_ptr, (int)M, (int)N, (int)K, stream); } else if (cfg == 1) { CUtensorMap a_map = g_maps.get(x, K, M, 128, 128, CuSw(128), CU_TENSOR_MAP_DATA_TYPE_UINT8, 1); CUtensorMap b_map = g_maps.get(w, K, N, 128, 256, CuSw(128), CU_TENSOR_MAP_DATA_TYPE_UINT8, 1); CUtensorMap c_map = g_maps.get((void*)y_ptr, N, M, 64, 128, CU_TENSOR_MAP_SWIZZLE_NONE, CU_TENSOR_MAP_DATA_TYPE_UINT16, 2); launch_cfg<128, 256, 128, 4, 4>(&a_map, &b_map, &c_map, (const float*)s_ptr, (int)M, (int)N, (int)K, stream); } else { CUtensorMap a_map = g_maps.get(x, K, M, 128, 128, CuSw(128), CU_TENSOR_MAP_DATA_TYPE_UINT8, 1); CUtensorMap b_map = g_maps.get(w, K, N, 128, 256, CuSw(128), CU_TENSOR_MAP_DATA_TYPE_UINT8, 1); CUtensorMap c_map = g_maps.get((void*)y_ptr, N, M, 64, 128, CU_TENSOR_MAP_SWIZZLE_NONE, CU_TENSOR_MAP_DATA_TYPE_UINT16, 2); launch_cfg<128, 256, 128, 4, 16>(&a_map, &b_map, &c_map, (const float*)s_ptr, (int)M, (int)N, (int)K, stream); } } void fp8_pad2(int64_t xa_ptr, int64_t xb_ptr, int64_t wa_ptr, int64_t wb_ptr, int64_t Ma, int64_t Mb, int64_t K, int64_t K_pad, int64_t stream_ptr) { const int dst_words = (int)(K_pad / 4); const int sa = (int)(Ma * K / 4); const int sb = (int)(Mb * K / 4); k_pad2_funnel<<<(int)(Ma + Mb), 256, 0, (cudaStream_t)stream_ptr>>>( (const uint32_t*)xa_ptr, (uint32_t*)xb_ptr, (const uint32_t*)wa_ptr, (uint32_t*)wb_ptr, (int)K, dst_words, sa, sb, (int)Ma); } void fp8_pad(int64_t src_ptr, int64_t dst_ptr, int64_t M, int64_t K, int64_t K_pad, int64_t stream_ptr) { const int dst_words = (int)(K_pad / 4); const int src_words = (int)(M * K / 4); k_pad_funnel<<<(int)M, 256, 0, (cudaStream_t)stream_ptr>>>( (const uint32_t*)src_ptr, (uint4*)dst_ptr, (int)K, dst_words, src_words); } """ _ext = load_inline( name="fp8_gemm_b200_v1", cpp_sources=[r""" void fp8_run(int64_t, int64_t, int64_t, int64_t, int64_t, int64_t, int64_t, int64_t, int64_t); void fp8_pad(int64_t, int64_t, int64_t, int64_t, int64_t, int64_t); void fp8_pad2(int64_t, int64_t, int64_t, int64_t, int64_t, int64_t, int64_t, int64_t, int64_t); """], cuda_sources=[_CUDA_SRC], functions=["fp8_run", "fp8_pad", "fp8_pad2"], extra_cuda_cflags=["-O3", "-std=c++17", "--use_fast_math", "-gencode=arch=compute_100a,code=sm_100a", "-lcuda"], extra_ldflags=["-lcuda"], verbose=False, ) def _raw_stream(): return torch._C._cuda_getCurrentRawStream(torch.cuda.current_device()) class Model(nn.Module): """y = ((x_fp8 @ w_fp8.T) * weight_scale).to(bf16) on tcgen05 tensor cores.""" def __init__(self, M: int, N: int, K: int): super().__init__() self.M, self.N, self.K = M, N, K w = torch.empty(N, K, dtype=torch.bfloat16) nn.init.normal_(w, std=0.02) s = (w.float().abs().amax(dim=1, keepdim=True) / E4M3_MAX).clamp(min=1e-12) w_fp8 = (w.float() / s).to(torch.float8_e4m3fn) self.register_buffer("weight", w_fp8) # (N, K) fp8 self.register_buffer("weight_scale", s.squeeze(1).to(torch.float32)) # (N,) self._run = _ext.fp8_run if M <= 64: self._cfg = 0 elif N >= 8192: self._cfg = 2 else: self._cfg = 1 # pad state for unaligned K self._w_pad = None # (version, buf) self._x_pad = None # buf for x def _pad(self, t, buf): _ext.fp8_pad(t.data_ptr(), buf.data_ptr(), t.shape[0], t.shape[1], buf.shape[1], _raw_stream()) return buf def forward(self, x: torch.Tensor) -> torch.Tensor: M, K = x.shape y = torch.empty((M, self.N), dtype=torch.bfloat16, device=x.device) stream = torch._C._cuda_getCurrentRawStream(0) if K % 16 == 0: w = self.weight K0 = K else: K0 = (K + 15) // 16 * 16 if self._x_pad is None: self._x_pad = torch.zeros((M, K0), dtype=x.dtype, device=x.device) w = self.weight ver = w._version need_w = self._w_pad is None or self._w_pad[0] != ver if need_w: wbuf = torch.zeros((self.N, K0), dtype=w.dtype, device=w.device) _ext.fp8_pad2(x.data_ptr(), self._x_pad.data_ptr(), w.data_ptr(), wbuf.data_ptr(), M, self.N, K, K0, stream) self._w_pad = (ver, wbuf) w = wbuf x = self._x_pad else: x = self._pad(x, self._x_pad) w = self._w_pad[1] self._run(x.data_ptr(), w.data_ptr(), self.weight_scale.data_ptr(), y.data_ptr(), M, self.N, K0, self._cfg, stream) return y Model.__call__ = Model.forward M = 4096 N = 4096 K = 4096 def get_inputs(): x = (torch.rand(M, K) * 8 - 4).to(torch.float8_e4m3fn) return [x] def get_init_inputs(): return [M, N, K]