submission 486943
dxyz · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 3661 lines, June 9 Researcher Reciprocity License v1.0.
submission_v1.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-nvfp4-group-gemm-486943?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:697d43eb7c40c18ae93b0ff6b37fa8822ffe815073e0ca6b3bfa99117ff0e215
license declaredunknown
license concludedunknown
authorsdxyz
imported2026-08-15
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
cluster
__cluster_dims__(CTA_GROUP, 1, 1)fp4
PyTorch reference implementation of NVFP4 block-scaled group GEMM.fused-epilogue
WaitEpilogue,mbarrier
void mbarrier_init(int mbar_addr, int count) {shared-memory
auto asmem_layout = make_layout(make_shape(_256{}, BLOCK_M, BLOCK_K / 256), make_stride(_1{}, _256{}, 256 * BLOCK_M)); // cute::LayoutLeft{} ?tcgen05
asm volatile("tcgen05.commit.cta_group::%2.mbarrier::arrive::one.shared::cluster.multicast::cluster.b64 [%0], %1;"tma
asm volatile("cp.async.bulk.shared::cluster.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({tmp[i * 4 + 0], tmp[i * 4 + 1]});Kernel source
submission_v1.py3661 lines
from task import input_t, output_t
from reference import ref_kernel
import torch
from torch.utils.cpp_extension import load_inline
COMMON_CU = r"""
#include <stdio.h>
#include <cudaTypedefs.h>
#include <cuda_fp16.h>
#include <torch/library.h>
#include <ATen/core/Tensor.h>
#include "cutlass/cutlass.h"
#include "cute/tensor.hpp"
constexpr int WARP_SIZE = 32;
constexpr int MMA_K = 64; // 32 bytes
constexpr int SF_BLOCK_SIZE = 16;
#define DIVUP(a, b) (((a) + (b) - 1) / (b))
// 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__
uint32_t elect_sync() {
uint32_t pred = 0;
asm volatile(
"{\n\t"
".reg .pred %%px;\n\t"
"elect.sync _|%%px, %1;\n\t"
"@%%px mov.s32 %0, 1;\n\t"
"}"
: "+r"(pred)
: "r"(0xFFFFFFFF)
);
return pred;
}
__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__
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)
);
}
template <int CTA_GROUP = 1>
__device__ inline
void tcgen05_commit_mcast(int mbar_addr, uint16_t cta_mask) {
asm volatile("tcgen05.commit.cta_group::%2.mbarrier::arrive::one.shared::cluster.multicast::cluster.b64 [%0], %1;"
:: "r"(mbar_addr), "h"(cta_mask), "n"(CTA_GROUP) : "memory");
}
// See https://docs.nvidia.com/cuda/parallel-thread-execution/#data-movement-and-conversion-instructions-bulk-copy
// cp.async.bulk is a non-blocking instruction which initiates an asynchronous bulk-copy operation
// from the location specified by source address operand srcMem
// to the location specified by destination address operand dstMem.
__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::cluster.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));
}
template <int CTA_GROUP = 1>
__device__ inline
void tma_1d_gmem2smem(int dst, const void *tmap_ptr, int x, int mbar_addr, uint64_t cache_policy) {
asm volatile("cp.async.bulk.tensor.1d.shared::cta.global.mbarrier::complete_tx::bytes.cta_group::%5.L2::cache_hint "
"[%0], [%1, {%2}], [%3], %4;"
:: "r"(dst), "l"(tmap_ptr), "r"(x), "r"(mbar_addr), "l"(cache_policy), "n"(CTA_GROUP)
: "memory");
}
template <int CTA_GROUP = 1>
__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::%7.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), "n"(CTA_GROUP)
: "memory");
}
// See https://docs.nvidia.com/cuda/parallel-thread-execution/#tcgen05-instructions-tcgen05-cp
// Instruction tcgen05.cp initiates an asynchronous copy operation from shared memory
// to the location specified by the address operand taddr in the Tensor Memory.
template <int CTA_GROUP = 1>
__device__ inline
void tcgen05_cp_nvfp4(int taddr, uint64_t s_desc) {
// .32x128b corresponds to lane x size. So we move 32 lanes and 128 bits = 128 / 8 bytes = 16 bytes which is 4 columns
// x4 means the data is multicast into 4 warps and each warp gets 1/4 of the data.
// .warpx4 populates data across 32-lane groups (lane in the sense of tmem).
// Some of the .shape qualifiers require certain .multicast qualifiers.
// .64x128b requires .warpx2::02_13 or .warpx2::01_23
// .32x128b requires .warpx4
//
// See https://docs.nvidia.com/cuda/parallel-thread-execution/#tcgen05-data-movement-shape
// 32x128b means 32 lanes, 128b == 16 bytes spans 16 / 4 == 4 columns
asm volatile("tcgen05.cp.cta_group::%2.32x128b.warpx4 [%0], %1;" :: "r"(taddr), "l"(s_desc), "n"(CTA_GROUP));
}
struct COLLECTOR_USAGE {
static constexpr char NONE[] = "";
static constexpr char A_FILL[] = ".collector::a::fill";
static constexpr char A_USE[] = ".collector::a::use";
static constexpr char A_LASTUSE[] = ".collector::a::lastuse";
static constexpr char A_DISCARD[] = ".collector::a::discard";
};
template <int CTA_GROUP = 1, const char *collector_usage = COLLECTOR_USAGE::NONE>
__device__ inline
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" // predicate register enable-input-d
"setp.ne.b32 p, %6, 0;\n\t"
"tcgen05.mma.cta_group::%7.kind::mxf4nvf4.block_scale.block16%8 [%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),
"n"(CTA_GROUP), "C"(collector_usage)
);
}
// 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); }
void check_cu(CUresult err) {
if (err == CUDA_SUCCESS) return;
const char *error_msg_ptr;
if (cuGetErrorString(err, &error_msg_ptr) != CUDA_SUCCESS)
error_msg_ptr = "unable to get error string";
TORCH_CHECK(false, "cuTensorMapEncodeTiled error: ", error_msg_ptr);
}
void check_cuda(cudaError_t err) {
if (err == cudaSuccess) return;
TORCH_CHECK(false, cudaGetErrorString(err));
}
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;
// Tile shape (boxDim) Global stride Comment
// [BLOCK_M, BLOCK_K] [K, 1] Original layout of threadblock tile in global memory.
// [BLOCK_M, BLOCK_K / 8, 8] [K, 8, 1] “Unflatten” the last dim.
// [BLOCK_K / 8, BLOCK_M, 8] [8, K, 1] Swap the first 2 dims.
// 8 if for bf16, 256 for nvfp4: we need contiguous strips of 256 elements
// The order needs to be reversed as cuTensorMapEncodeTiled assumes the tensors are
// col major.
// For the smem layout see https://docs.nvidia.com/cuda/parallel-thread-execution/#tcgen05-shared-memory-layout-swizzling
// See also https://docs.nvidia.com/cuda/parallel-thread-execution/#asynchronous-warpgroup-level-matrix-shared-memory-layout
//
// For example, 128B MN major swizzle atom would have a shape of (8*(128/32)) x 8 = 32x8 for tf32 tensor core inputs.
// So for 4 bits we have (8 * 128 / 4) x 8 = 256 x 8
// the swizzling pattern is given in the tables.
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}; // When all elements of elementStrides array is one, boxDim specifies the number of elements to load.
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);
}
void init_SF_tmap(CUtensorMap *tmap, const char *ptr, uint64_t global_size, uint32_t shared_size) {
// use int64 as dtype, hence divide sizes by 8
constexpr uint32_t rank = 1;
uint64_t globalDim[rank] = {global_size / 8};
uint64_t globalStrides[rank-1] = {}; // in bytes
uint32_t boxDim[rank] = {shared_size / 8};
uint32_t elementStrides[rank] = {1};
auto err = cuTensorMapEncodeTiled(
tmap,
CUtensorMapDataType::CU_TENSOR_MAP_DATA_TYPE_INT64,
rank,
(void *)ptr,
globalDim,
globalStrides,
boxDim,
elementStrides,
CUtensorMapInterleave::CU_TENSOR_MAP_INTERLEAVE_NONE,
CUtensorMapSwizzle::CU_TENSOR_MAP_SWIZZLE_NONE,
CUtensorMapL2promotion::CU_TENSOR_MAP_L2_PROMOTION_NONE,
CUtensorMapFloatOOBfill::CU_TENSOR_MAP_FLOAT_OOB_FILL_NONE
);
//check_cu(err);
}
"""
CUDA_SRC_0 = r"""
using namespace cute;
constexpr int MAX_TMEM_COLS = 512;
constexpr int NUM_SMS = 148;
constexpr int NUM_BLOCKS0 = 64;
constexpr int NUM_BLOCKS1 = 96;
constexpr int NUM_BLOCKS2 = 64;
constexpr int NUM_BLOCKS3 = 64;
constexpr int NUM_BLOCKS4 = 32;
constexpr int NUM_BLOCKS5 = 128;
constexpr int NUM_BLOCKS6 = 64;
constexpr int NUM_BLOCKS7 = 96;
constexpr int GEMM_BLOCK_END0 = NUM_BLOCKS0;
constexpr int GEMM_BLOCK_END1 = GEMM_BLOCK_END0 + NUM_BLOCKS1;
constexpr int GEMM_BLOCK_END2 = GEMM_BLOCK_END1 + NUM_BLOCKS2;
constexpr int GEMM_BLOCK_END3 = GEMM_BLOCK_END2 + NUM_BLOCKS3;
constexpr int GEMM_BLOCK_END4 = GEMM_BLOCK_END3 + NUM_BLOCKS4;
constexpr int GEMM_BLOCK_END5 = GEMM_BLOCK_END4 + NUM_BLOCKS5;
constexpr int GEMM_BLOCK_END6 = GEMM_BLOCK_END5 + NUM_BLOCKS6;
constexpr int GEMM_BLOCK_END7 = GEMM_BLOCK_END6 + NUM_BLOCKS7;
constexpr int NUM_BLOCKS = GEMM_BLOCK_END7;
template<int PROBLEM_SIZE>
__constant__ CUtensorMap A_tmaps[PROBLEM_SIZE];
template<int PROBLEM_SIZE>
__constant__ CUtensorMap B_tmaps[PROBLEM_SIZE];
template<int PROBLEM_SIZE>
__constant__ CUtensorMap SFA_tmaps[PROBLEM_SIZE];
template<int PROBLEM_SIZE>
__constant__ CUtensorMap SFB_tmaps[PROBLEM_SIZE];
template<int PROBLEM_SIZE>
__constant__ half* C_ptrs[PROBLEM_SIZE];
template<int PROBLEM_SIZE>
__constant__ int Ms[PROBLEM_SIZE];
template<int PROBLEM_SIZE>
__constant__ int Ns[PROBLEM_SIZE];
template <
int K,
int BLOCK_M,
int BLOCK_N,
int BLOCK_K,
int NUM_STAGES,
class SFALayout,
class SFBLayout,
bool SWAP_AB,
int PROBLEM_SIZE,
int CTA_GROUP
>
__global__
__cluster_dims__(CTA_GROUP, 1, 1)
__launch_bounds__(BLOCK_M + 2 * WARP_SIZE) // __launch_bounds__(MAX_THREADS_PER_BLOCK, MIN_BLOCKS_PER_MP)
void multi_gemm_kernel_p0() {
// https://docs.nvidia.com/cuda/parallel-thread-execution/#tcgen05-matrix-shape
static_assert(BLOCK_M == 128); // There doesnt seem to be any code for tiling block_M
// This layout is used as TMA assumes col major, this is a colasced canonical layout.
auto asmem_layout = make_layout(make_shape(_256{}, BLOCK_M, BLOCK_K / 256), make_stride(_1{}, _256{}, 256 * BLOCK_M)); // cute::LayoutLeft{} ?
// Each member of the cluster will get a cta_rank specific offset
auto bsmem_layout = make_layout(make_shape(_256{}, BLOCK_N / CTA_GROUP, BLOCK_K / 256), make_stride(_1{}, _256{}, 256 * BLOCK_N / CTA_GROUP));
// an SF atom is ((32, 4), (16, 4)), the shmem layout is ((128 * SF_ATOMS_PER_BLOCK_K, 4)) for TMA purposes
// Since we are using CU_TENSOR_MAP_SWIZZLE_NONE, the TMA atom has shape (8, 128bits) or (8, 16) bytes, so SBO below is 8 * 16
constexpr int SF_TILE_ROWS = 128;
constexpr int SF_ATOM_NUM_ELEMENTS = 128 * 4; // cosize(SF_atom)
constexpr int SF_ATOMS_PER_BLOCK_K = BLOCK_K / 16 / 4; // each atom has 16 * 4 columns
constexpr int SF_ROWS_PER_TMEM_COL = 32; // each tmem col has 32 rows - since we are using tcgen05.cp.cta_group::1.32x128b.warpx4
using Problem = std::tuple<const CUtensorMap*, const CUtensorMap*, const CUtensorMap*, const CUtensorMap*, half*, int, int, SFALayout, SFBLayout>;
auto get_problem = [&](int problem_id) -> Problem {
auto make_sf_layout = [](int rest) {
constexpr int rest_K = K / 16 / 4;
return make_layout(make_shape(make_shape(make_shape(_32{}, _4{}), rest), make_shape(make_shape(_16{}, _4{}), rest_K)),
make_stride(make_stride(make_stride(_16{}, _4{}), rest_K * 512), make_stride(make_stride(_0{}, _1{}), _512{})));
};
const CUtensorMap* A_tmap = &A_tmaps<PROBLEM_SIZE>[problem_id];
const CUtensorMap* B_tmap = &B_tmaps<PROBLEM_SIZE>[problem_id];
const CUtensorMap* SFA_tmap = &SFA_tmaps<PROBLEM_SIZE>[problem_id];
const CUtensorMap* SFB_tmap = &SFB_tmaps<PROBLEM_SIZE>[problem_id];
half* C = C_ptrs<PROBLEM_SIZE>[problem_id];
int M = Ms<PROBLEM_SIZE>[problem_id];
int N = Ns<PROBLEM_SIZE>[problem_id];
const int rest_M = DIVUP(M, SF_TILE_ROWS);
const int rest_N = DIVUP(N, SF_TILE_ROWS);
SFALayout sfa_layout = make_sf_layout(rest_M);
SFBLayout sfb_layout = make_sf_layout(rest_N);
return {A_tmap, B_tmap, SFA_tmap, SFB_tmap, C, M, N, sfa_layout, sfb_layout};
};
// CTA rank in a cluster
int cta_rank;
asm volatile("mov.b32 %0, %%cluster_ctarank;" : "=r"(cta_rank));
const int tid = threadIdx.x;
const int lane_id = tid % WARP_SIZE;
const int warp_id = tid / WARP_SIZE;
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; // BLOCK_M * BLOCK_K elements, each element is 4 bits so divide by 2 to get bytes.
constexpr int B_size = BLOCK_N * BLOCK_K / 2 / CTA_GROUP; // each CTA only loads half of B
constexpr int SFA_size = SF_ATOM_NUM_ELEMENTS * SF_ATOMS_PER_BLOCK_K; // always copy one atom - this is ATOM_SIZE_IN_BYTES * ATOMS_PER_BLOCK_K
constexpr int SFB_size = SF_ATOM_NUM_ELEMENTS * SF_ATOMS_PER_BLOCK_K;
constexpr int STAGE_SIZE = A_size + B_size + SFA_size + SFB_size;
// set up mbarriers and tmem
// we have NUM_STAGES mbars for TMA
// NUM_STAGES mbars for MMA
// 1 mbar for mainloop
#pragma nv_diag_suppress static_var_with_dynamic_init
__shared__ int64_t mbars[NUM_STAGES * 2 + 1];
__shared__ int tmem_addr[1]; // tmem address is 32-bit
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;
const int epilogue_mbar_addr = mainloop_mbar_addr + 2 * 8;
// https://docs.nvidia.com/cuda/parallel-thread-execution/#tcgen05-mma-scale-factor-a-layout-4x
// each MMA consumes:
// - (128, 64) of A -> (128, 4) of SFA -> reshaped as (32, 4', 4) -> 4 tmem columns
constexpr int SFA_tmem = BLOCK_N * 2; // Double buffer the mma output.
constexpr int SFB_tmem = SFA_tmem + 4 * (BLOCK_K / MMA_K); // we need BLOCK_K / MMA_K many scale factors for the "tiledMMA"
if (warp_id == 0 && elect_sync()) {
for (int prefetch_idx = 0; prefetch_idx < 8; ++prefetch_idx) {
asm volatile("prefetch.tensormap [%0];" :: "l"(&A_tmaps<PROBLEM_SIZE>[prefetch_idx]) : "memory");
asm volatile("prefetch.tensormap [%0];" :: "l"(&B_tmaps<PROBLEM_SIZE>[prefetch_idx]) : "memory");
asm volatile("prefetch.tensormap [%0];" :: "l"(&SFA_tmaps<PROBLEM_SIZE>[prefetch_idx]) : "memory");
asm volatile("prefetch.tensormap [%0];" :: "l"(&SFB_tmaps<PROBLEM_SIZE>[prefetch_idx]) : "memory");
}
} else if (warp_id == 1 && elect_sync()) {
for (int i = 0; i < NUM_STAGES; i++) {
mbarrier_init(tma_mbar_addr + i * 8, CTA_GROUP); // one thread in each cluster CTA reports.
mbarrier_init(mma_mbar_addr + i * 8, 1);
}
mbarrier_init(mainloop_mbar_addr, 1); // barrier for first tmem mma buffer
mbarrier_init(mainloop_mbar_addr + 8, 1); // barrier for second tmem mma buffer
mbarrier_init(epilogue_mbar_addr, 4 * CTA_GROUP);
mbarrier_init(epilogue_mbar_addr + 8, 4 * CTA_GROUP);
asm volatile("fence.mbarrier_init.release.cluster;"); // visible to async proxy
}
if constexpr (CTA_GROUP > 1) {
// visible to all threads in a cluster
asm volatile("barrier.cluster.arrive.relaxed.aligned;");
asm volatile("barrier.cluster.wait.acquire.aligned;");
}
else {
// visible to all threads in a threadblock
__syncthreads();
}
constexpr int num_iters = K / BLOCK_K;
const int bid = blockIdx.x;
auto scheduler = [&](int global_cluster_id) -> std::tuple<int, int, int, int> {
int problem_id = 0;
int local_cluster_id = global_cluster_id;
if (global_cluster_id < GEMM_BLOCK_END0 / CTA_GROUP) {
problem_id = 0;
local_cluster_id = global_cluster_id;
} else if (global_cluster_id < GEMM_BLOCK_END1 / CTA_GROUP) {
problem_id = 1;
local_cluster_id = global_cluster_id - GEMM_BLOCK_END0 / CTA_GROUP;
} else if (global_cluster_id < GEMM_BLOCK_END2 / CTA_GROUP) {
problem_id = 2;
local_cluster_id = global_cluster_id - GEMM_BLOCK_END1 / CTA_GROUP;
} else if (global_cluster_id < GEMM_BLOCK_END3 / CTA_GROUP) {
problem_id = 3;
local_cluster_id = global_cluster_id - GEMM_BLOCK_END2 / CTA_GROUP;
} else if (global_cluster_id < GEMM_BLOCK_END4 / CTA_GROUP) {
problem_id = 4;
local_cluster_id = global_cluster_id - GEMM_BLOCK_END3 / CTA_GROUP;
} else if (global_cluster_id < GEMM_BLOCK_END5 / CTA_GROUP) {
problem_id = 5;
local_cluster_id = global_cluster_id - GEMM_BLOCK_END4 / CTA_GROUP;
} else if (global_cluster_id < GEMM_BLOCK_END6 / CTA_GROUP) {
problem_id = 6;
local_cluster_id = global_cluster_id - GEMM_BLOCK_END5 / CTA_GROUP;
} else if (global_cluster_id < GEMM_BLOCK_END7 / CTA_GROUP) {
problem_id = 7;
local_cluster_id = global_cluster_id - GEMM_BLOCK_END6 / CTA_GROUP;
}
int N = Ns<PROBLEM_SIZE>[problem_id];
const int grid_n = DIVUP(N, BLOCK_N) ;
// The cluster Ids are arranged in a (grid_m / 2 , grid_n) matrix in a row major layout.
const int bid_m = local_cluster_id / grid_n * CTA_GROUP + cta_rank;
const int bid_n = local_cluster_id % grid_n;
const int off_m = bid_m * BLOCK_M;
const int cluster_off_n = bid_n * BLOCK_N; // N offset for this cluster
const int off_n = cluster_off_n + cta_rank * (BLOCK_N / CTA_GROUP); // N offset for this CTA
return {off_m, off_n, cluster_off_n, problem_id};
};
// warp-specialization
if (warp_id == NUM_WARPS - 2 && elect_sync()) {
// TMA warp
auto issue_tma = [&](int iter_k, int stage_id, int off_m, int off_n, int problem_id) {
const int mbar_addr = (tma_mbar_addr + stage_id * 8) & 0xFEFFFFFF; // CTA0's barrier
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 problem = get_problem(problem_id);
const CUtensorMap* A_tmap = std::get<0>(problem);
const CUtensorMap* B_tmap = std::get<1>(problem);
const int M = std::get<5>(problem);
const int N = std::get<6>(problem);
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;
}
// issue TMA
const int off_k = iter_k * BLOCK_K;
tma_3d_gmem2smem<CTA_GROUP>(A_smem, A_tmap, 0, off_m, off_k / 256, mbar_addr, cache_A);
tma_3d_gmem2smem<CTA_GROUP>(B_smem, B_tmap, 0, off_n, off_k / 256, mbar_addr, cache_B);
const int rest_m = off_m / 32 / 4; // M atom is (32, 4)
const int rest_n = off_n / 32 / 4; // N atom is (32, 4)
const int rest_k = off_k / 16 / 4; // The SF atom is ((32, 4), (16, 4))
const CUtensorMap* SFA_tmap = std::get<2>(problem);
const CUtensorMap* SFB_tmap = std::get<3>(problem);
SFALayout sfa_layout = std::get<7>(problem);
SFALayout sfb_layout = std::get<8>(problem);
// Divide by 8 since underlying type is INT64
int sfa_offset = crd2idx(make_coord(make_coord(_0{}, rest_m), make_coord(_0{}, rest_k)), sfa_layout) / 8;
tma_1d_gmem2smem<CTA_GROUP>(SFA_smem, SFA_tmap, sfa_offset, mbar_addr, cache_A);
int sfb_offset = crd2idx(make_coord(make_coord(_0{}, rest_n), make_coord(_0{}, rest_k)), sfb_layout) / 8;
tma_1d_gmem2smem<CTA_GROUP>(SFB_smem, SFB_tmap, sfb_offset, mbar_addr, cache_B);
// signal TMA done .shared::cluster
asm volatile("mbarrier.arrive.expect_tx.release.cta.shared::cluster.b64 _, [%0], %1;"
:: "r"(mbar_addr), "r"(STAGE_SIZE) : "memory");
};
int mma_phase = 1; // Start with phase 1 as there are no pending MMA operations
int tma_pipeline_stage = 0;
for (int this_bid = bid; this_bid < NUM_BLOCKS; this_bid += NUM_SMS) {
auto [off_m, off_n, dummy, problem_id] = scheduler(this_bid / CTA_GROUP); // divide by CTA_GROUP to get cluster id
for (int iter_k = 0; iter_k < num_iters; iter_k++) {
// wait MMA
// NB when debugging, it will crash w/o this barrier. You cannot keep issuing TMAs w/o the
// signal the mbarrier_wait(tma_mbar_addr + stage_id * 8, tma_phase); has completed.
mbarrier_wait(mma_mbar_addr + tma_pipeline_stage * 8, mma_phase);
issue_tma(iter_k, tma_pipeline_stage, off_m, off_n, problem_id);
tma_pipeline_stage = (tma_pipeline_stage + 1) % NUM_STAGES;
if (tma_pipeline_stage == 0) {
mma_phase ^= 1;
}
}
}
}
else if (warp_id == NUM_WARPS - 1) {
// allocate tmem
const int addr = static_cast<int>(__cvta_generic_to_shared(tmem_addr));
asm volatile("tcgen05.alloc.cta_group::%2.sync.aligned.shared::cta.b32 [%0], %1;" :: "r"(addr), "r"(MAX_TMEM_COLS), "n"(CTA_GROUP));
if (cta_rank == 0 && elect_sync()) {
// MMA warp
// https://docs.nvidia.com/cuda/parallel-thread-execution/#tcgen05-instruction-descriptor
// fp4 MMA doesn't support MMA_M=64. Hence, we will use MMA_M=128 and ignore the rest.
constexpr int MMA_N = BLOCK_N;
constexpr int MMA_M = 128 * CTA_GROUP;
constexpr int16_t cta_mask = (1 << CTA_GROUP) - 1;
constexpr uint32_t i_desc = (1U << 7U) // atype=E2M1
| (1U << 10U) // btype=E2M1
| ((uint32_t)MMA_N >> 3U << 17U)
| ((uint32_t)MMA_M >> 7U << 27U)
;
int tma_phase = 0;
int mma_pipeline_stage = 0;
int epilogue_phase = 1;
int mainloop_stage = 0; // 0/1 indicates where in tmem we will put the result
for (int this_bid = bid; this_bid < NUM_BLOCKS; this_bid += NUM_SMS) {
mbarrier_wait(epilogue_mbar_addr + 8 * mainloop_stage, epilogue_phase);
auto [off_m, off_n, dummy, dummy1] = scheduler(this_bid / CTA_GROUP); // divide by CTA_GROUP to get cluster id
// Given a k slice of the scale factors which is an atom, the scale factors are stored in groups of four columns in tmem
// These will be aligned to rows 0, SF_TILE_ROWS, SF_TILE_ROWS * 2, ...
// And the sf 0..31 will be in index_col=0, 32..63 index_col =1 etc
// Given an arbitrary row_offset in M or N, in order to find the tmem column offset,
// Find the residue at the SF tile row granularity, (row_offset % SF_TILE_ROWS) and
// since each tmem col stores 32 SF, move over by the desired number of columns
// sf_tmem_offset = (row_offset % SF_TILE_ROWS) / SF_ROWS_PER_TMEM_COL;
const int M_sf_tmem_offset = (off_m % SF_TILE_ROWS) / SF_ROWS_PER_TMEM_COL;
const int N_sf_tmem_offset = (off_n % SF_TILE_ROWS) / SF_ROWS_PER_TMEM_COL;
for (int iter_k = 0; iter_k < num_iters; iter_k++) {
// wait TMA
mbarrier_wait(tma_mbar_addr + mma_pipeline_stage * 8, tma_phase);
const int A_smem = smem + mma_pipeline_stage * 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;
// 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.
auto make_desc_AB = [](int addr) -> uint64_t {
// GMMA::Layout_K_SW128_Atom<cutlass::float_e2m1_t>{} is (_8,_256):(_256,_1) so to get to the start of the next row is 256 / 2 = 128 bytes
// so to skip 8 rows, we need 8 * 128 bytes.
const int SBO = 8 * 128;
return desc_encode(addr) | (desc_encode(SBO) << 32ULL) | (1ULL << 46ULL) | (2ULL << 61ULL);
};
// no swizzling
auto make_desc_SF = [](int addr) -> uint64_t {
// Since we are using no swizzling, the atom has size (8, 16 = 128 / 8). These atoms
// are tiled (SF_ATOMS_PER_BLOCK_K = 4, 1). So the SBO is 8 * 16
const int SBO = 8 * 16;
return desc_encode(addr) | (desc_encode(SBO) << 32ULL) | (1ULL << 46ULL);
};
// tcgen05.cp -> tcgen05.mma should be pipelined correctly per PTX doc
// https://docs.nvidia.com/cuda/parallel-thread-execution/#tcgen05-memory-consistency-model-pipelined-instructions
// cutlass issues all of smem->tmem BEFORE mma
// https://github.com/NVIDIA/cutlass/blob/v4.3.2/include/cutlass/gemm/collective/sm100_blockscaled_mma_warpspecialized.hpp#L1013-L1016
constexpr uint64_t SF_desc = make_desc_SF(0);
// The below should really be SF_desc + matrix-descriptor-encode(x) , where matrix-descriptor-encode(x) = (x & 0x3FFFF) >> 4
// but (x & 0x3FFFF) >> 4 simplifies to x >> 4.
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++) {
// The below should really be SFA_desc + (uint64_t)k * matrix-descriptor-encode(512ULL); where matrix-descriptor-encode(x) = (x & 0x3FFFF) >> 4
// but (x & 0x3FFFF) >> 4 simplifies to x >> 4. So this is SFA_desc + (uint64_t)k * (512ULL >> 4ULL)
// So we are copying 512 bytes chunks , so adjust the shmem desc accordingly.
uint64_t sfa_desc = SFA_desc + (uint64_t)k * (512ULL >> 4ULL); // 4 columns, 512 bytes of 128x4 / 32x4x4
uint64_t sfb_desc = SFB_desc + (uint64_t)k * (512ULL >> 4ULL);
// tmem addresses
// Tensor Memory addresses are 32-bit wide and specify two components.
// Lane index
// Column index
// The layout is as follows:
// 31 16
// 15 0
// Lane index
// Column index
// We are using tcgen05.cp.cta_group::1.32x128b.warpx4
// See https://docs.nvidia.com/cuda/parallel-thread-execution/#tcgen05-mma-scale-factor-a-layout-4x
// so the slices ((32, 4), (16, 4)) slices of the scale factors are in the 32 rows (lanes) of tmem
tcgen05_cp_nvfp4<CTA_GROUP>(SFA_tmem + k * 4, sfa_desc);
tcgen05_cp_nvfp4<CTA_GROUP>(SFB_tmem + k * 4, sfb_desc);
}
// Compare to tiled_mma
// k1 selects the (BLOCK_M, 256) tile.
// k2 selects the (BLOCK_M, 64) tile, whose rows are swizzled.
// NOTE: this doesn't work with BLOCK_N=32, since apparently tcgen05.mma requires SFB_tmem
// to have 2-column (8-byte) alignment (looks like not documented).
// HOWEVER: works with a cluster with BLOCK_N = 64 - maybe only the issuing thread has the alignment
// requirement.
for (int k1 = 0; k1 < BLOCK_K / 256; k1++) { // BLOCK_K == 256 so this loop only executes once
for (int k2 = 0; k2 < 256 / MMA_K; k2++) { // this is the inner loop needed since MMA_K is in general < 256 = BLOCK_K
// The layout in shmem is the result of the TMA: {256, BLOCK_M, BLOCK_K / 256};
// Notice that in TMA, it’s implied that the shared memory destination has the natural contiguous layout of the given shape
// (we don’t specify shared memory stride anywhere).
// since cuTensorMapEncodeTiled assumes the tensors are col major cute::LayoutLeft
// The layout is then (256, BLOCK_M, BLOCK_K / 256) : (1, 256, 256 * BLOCK_M) elements
// https://gau-nernst.github.io/tcgen05/#decipher-tcgen05
//
// Thus given k1 and k2, the coords are (k2 * MMA_K, 0, k1)
// so the offset in elements is k1 * 256 * BLOCK_M + k2 * MMA_K = k1 * 256 * BLOCK_M + k2 * 64 (MMA_K = 64)
// so the offset in bytes = k1 * 128 * BLOCK_M + k2 * 32
// which is the expression below.
// uint64_t a_desc = make_desc_AB(A_smem + k1 * BLOCK_M * 128 + k2 * 32);
// crd2idx(make_coord(k2 * MMA_K, 0, k1), asmem_layout) is in elements, divide by 2 to get bytes
uint64_t a_desc = make_desc_AB(A_smem + crd2idx(make_coord(k2 * MMA_K, _0{}, k1), asmem_layout) / 2);
uint64_t b_desc = make_desc_AB(B_smem + crd2idx(make_coord(k2 * MMA_K, _0{}, k1), bsmem_layout) / 2);
// tmem addresses
// Tensor Memory addresses are 32-bit wide and specify two components.
// Lane index
// Column index
// The layout is as follows:
// 31 16
// 15 0
// Lane index
// Column index
int k_sf = k1 * 4 + k2; // the scale factors are in blocks of 4. This is the block idx
// k_sf is mulitplied by 4 to get us to the actual column.
const int scale_A_tmem = SFA_tmem + k_sf * 4 + M_sf_tmem_offset;
const int scale_B_tmem = SFB_tmem + k_sf * 4 + N_sf_tmem_offset;
const int enable_input_d = (k1 == 0 && k2 == 0) ? iter_k : 1;
tcgen05_mma_nvfp4<CTA_GROUP>(/*dst tmem*/ BLOCK_N * mainloop_stage, a_desc, b_desc, i_desc, scale_A_tmem, scale_B_tmem, enable_input_d);
} // for k2
} // for k1
// signal MMA done
// @@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@
// tcgen05_commit_mcast does not work if CTA_GROUP = 1
// tcgen05_commit_mcast<CTA_GROUP>(mma_mbar_addr + mma_pipeline_stage * 8, cta_mask); // signal MMA done
// @@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@
asm volatile("tcgen05.commit.cta_group::%2.mbarrier::arrive::one.shared::cluster.multicast::cluster.b64 [%0], %1;"
:: "r"(mma_mbar_addr + mma_pipeline_stage * 8), "h"(cta_mask), "n"(CTA_GROUP) : "memory");
mma_pipeline_stage = (mma_pipeline_stage + 1) % NUM_STAGES;
if (mma_pipeline_stage == 0) {
tma_phase ^= 1;
}
}
// signal mainloop done
// tcgen05_commit_mcast<CTA_GROUP>(mainloop_mbar_addr + mainloop_stage * 8, cta_mask); // signal mainloop done
asm volatile("tcgen05.commit.cta_group::%2.mbarrier::arrive::one.shared::cluster.multicast::cluster.b64 [%0], %1;"
:: "r"(mainloop_mbar_addr + mainloop_stage * 8), "h"(cta_mask), "n"(CTA_GROUP) : "memory");
mainloop_stage = (mainloop_stage + 1) % 2;
if (mainloop_stage == 0) {
epilogue_phase ^= 1;
}
} // for bid
} // cta_rank == 0 && elect_sync()
} // warp_id == NUM_WARPS - 1
else if (tid < BLOCK_M) {
// epilogue warps
int mainloop_phase = 0;
int mainloop_stage = 0; // 0/1 indicates where in tmem we will put the result
for (int this_bid = bid; this_bid < NUM_BLOCKS; this_bid += NUM_SMS) {
auto result = scheduler(this_bid / CTA_GROUP); // divide by CTA_GROUP to get cluster id
int off_m = std::get<0>(result);
int cluster_off_n = std::get<2>(result);
int problem_id = std::get<3>(result);
auto problem = get_problem(problem_id);
const int M = std::get<5>(problem);
const int N = std::get<6>(problem);
half* C_ptr = std::get<4>(problem);
// wait mainloop
mbarrier_wait(mainloop_mbar_addr + mainloop_stage * 8, mainloop_phase);
asm volatile("tcgen05.fence::after_thread_sync;");
int tmem_buffer_offset = BLOCK_N * mainloop_stage;
auto epilogue_M_major = [&]() {
// C is M-major
constexpr int WIDTH = std::min(BLOCK_N, 64); // using 128 might be slower
for (int n = 0; n < (BLOCK_N + WIDTH - 1)/ WIDTH; n++) {
float tmp[WIDTH]; // if WIDTH=128, we are using 128 registers here
if constexpr (WIDTH == 128) tcgen05_ld_32x32bx128(tmp, warp_id * 32, n * WIDTH + tmem_buffer_offset);
if constexpr (WIDTH == 64) tcgen05_ld_32x32bx64(tmp, warp_id * 32, n * WIDTH + tmem_buffer_offset);
if constexpr (WIDTH == 32) tcgen05_ld_32x32bx32(tmp, warp_id * 32, n * WIDTH + tmem_buffer_offset);
asm volatile("tcgen05.wait::ld.sync.aligned;");
for (int i = 0; i < WIDTH; i++) {
// NB row and col are transposed.
const int row = cluster_off_n + n * WIDTH + i;
const int col = off_m + tid;
if (row < N && col < M) {
C_ptr[row * M + col] = __float2half(tmp[i]);
}
}
}
};
auto epilogue_N_major = [&]() {
// C is N-major
for (int m = 0; m < 32 / 16; m++) {
float tmp[BLOCK_N / 2];
if constexpr (BLOCK_N == 128) tcgen05_ld_16x256bx16(tmp, warp_id * 32 + m * 16, tmem_buffer_offset);
if constexpr (BLOCK_N == 64) tcgen05_ld_16x256bx8(tmp, warp_id * 32 + m * 16, tmem_buffer_offset);
if constexpr (BLOCK_N == 32) tcgen05_ld_16x256bx4(tmp, warp_id * 32 + m * 16, tmem_buffer_offset);
asm volatile("tcgen05.wait::ld.sync.aligned;");
for (int i = 0; i < (BLOCK_N + 7)/ 8; i++) {
// TODO replace this with cute tensors.
const int row = off_m + warp_id * 32 + m * 16 + lane_id / 4;
const int col = cluster_off_n + i * 8 + (lane_id % 4) * 2;
if (row < M) {
reinterpret_cast<half2 *>(C_ptr + (row + 0) * N + col)[0] = __float22half2_rn({tmp[i * 4 + 0], tmp[i * 4 + 1]});
}
if (row + 8 < M) {
reinterpret_cast<half2 *>(C_ptr + (row + 8) * N + col)[0] = __float22half2_rn({tmp[i * 4 + 2], tmp[i * 4 + 3]});
}
}
}
};
if constexpr (SWAP_AB) {
epilogue_M_major();
} else {
epilogue_N_major();
}
if (elect_sync()) {
// // signal Epilogue done
// constexpr int16_t cta_mask = 3;
// tcgen05_commit_mcast<CTA_GROUP>(epilogue_mbar_addr + 8 * mainloop_stage, cta_mask); // signal Epilogue done
const int mbar_addr = (epilogue_mbar_addr + mainloop_stage * 8) & 0xFEFFFFFF;
asm volatile("mbarrier.arrive.release.cta.shared::cluster.b64 _, [%0];" :: "r"(mbar_addr) : "memory");
}
mainloop_stage = (mainloop_stage + 1) % 2;
if (mainloop_stage == 0) {
mainloop_phase ^= 1;
}
} // for bid
} //epilogue warps
// All warps.
if constexpr (CTA_GROUP > 1) {
// Dont exit w/o the peer CTA.
asm volatile("barrier.cluster.arrive.release.aligned;");
asm volatile("barrier.cluster.wait.acquire.aligned;");
} else {
__syncthreads(); // all threads finish reading data from tmem
}
// deallocate tmem. tmem address should be 0.
if (warp_id == 0) {
asm volatile("tcgen05.dealloc.cta_group::%2.sync.aligned.b32 %0, %1;" :: "r"(0), "r"(MAX_TMEM_COLS), "n"(CTA_GROUP));
}
}
template <
int K,
int BLOCK_M,
int BLOCK_N,
int BLOCK_K,
int NUM_STAGES,
bool SWAP_AB
>
std::vector<at::Tensor> launch_gemm_cstr_p0(
c10::List<std::tuple<at::Tensor, at::Tensor, at::Tensor>> ABC,
c10::List<std::tuple<at::Tensor, at::Tensor>> SFAB,
c10::List<at::Tensor> outputs) {
constexpr int CTA_GROUP = 2;
static_assert(BLOCK_K % 256 == 0);
constexpr int rest_K = K / 16 / 4;
constexpr int grid = NUM_SMS;
constexpr int tb_size = BLOCK_M + 2 * WARP_SIZE;
constexpr int AB_size = (BLOCK_M + BLOCK_N / CTA_GROUP) * (BLOCK_K / 2);
constexpr int SFAB_size = 128 * (BLOCK_K / 16) * 2;
constexpr int smem_size = (AB_size + SFAB_size) * NUM_STAGES;
// This layout is used as TMA assumes col major, this is a colasced canonical layout.
auto layout_smem_A = make_layout(make_shape(256, BLOCK_M, BLOCK_K / 256), make_stride(1, 256, 256 * BLOCK_M)); // cute::LayoutLeft{} ?
// Each member of the cluster will get a cta_rank specific offset
auto layout_smem_B = make_layout(make_shape(256, BLOCK_N / CTA_GROUP, BLOCK_K / 256), make_stride(1, 256, 256 * BLOCK_N / CTA_GROUP));
constexpr int SF_TILE_ROWS = 128;
constexpr int SF_ATOM_NUM_ELEMENTS = 128 * 4; // cosize(SF_atom)
constexpr int SF_ATOMS_PER_BLOCK_K = BLOCK_K / 16 / 4; // each atom has 16 * 4 columns
auto make_sf_layout = [&](int rest) {
return make_layout(make_shape(make_shape(make_shape(_32{}, _4{}), rest), make_shape(make_shape(_16{}, _4{}), rest_K)),
make_stride(make_stride(make_stride(_16{}, _4{}), rest_K * 512), make_stride(make_stride(_0{}, _1{}), _512{})));
};
// This is for the type only
auto layout_sfa = make_sf_layout(16);
auto layout_sfb = make_sf_layout(16);
auto populate_kernel_config = [&ABC, &SFAB, &outputs](CUtensorMap* A_tmaps, CUtensorMap* B_tmaps, CUtensorMap* SFA_tmaps, CUtensorMap* SFB_tmaps,
half** C_ptrs, int* Ms, int* Ns, auto problem_size_const) {
constexpr int PROBLEM_SIZE = decltype(problem_size_const)::value;
for (int workItem = 0; workItem < PROBLEM_SIZE; ++workItem) {
std::tuple<at::Tensor, at::Tensor, at::Tensor> abc_tuple = ABC[workItem];
auto A = std::get<0>(abc_tuple);
auto B = std::get<1>(abc_tuple);
std::tuple<at::Tensor, at::Tensor> sfab_tuple = SFAB[workItem];
auto SFA = std::get<0>(sfab_tuple);
auto SFB = std::get<1>(sfab_tuple);
at::Tensor C = outputs[workItem];
const int M = A.size(0);
const int N = B.size(0);
auto A_ptr = reinterpret_cast<const char *>(A.data_ptr());
auto B_ptr = reinterpret_cast<const char *>(B.data_ptr());
auto SFA_ptr = reinterpret_cast<const char *>(SFA.data_ptr());
auto SFB_ptr = reinterpret_cast<const char *>(SFB.data_ptr());
C_ptrs[workItem] = reinterpret_cast<half *>(C.data_ptr());
int new_M = M;
int new_N = N;
if constexpr (SWAP_AB) {
std::swap(A_ptr, B_ptr);
std::swap(SFA_ptr, SFB_ptr);
std::swap(new_M, new_N);
}
Ms[workItem] = new_M;
Ns[workItem] = new_N;
init_AB_tmap(&A_tmaps[workItem], A_ptr, new_M, K, BLOCK_M, BLOCK_K);
init_AB_tmap(&B_tmaps[workItem], B_ptr, new_N, K, BLOCK_N / CTA_GROUP, BLOCK_K);
// CUtensorMap SFA_tmap0, SFB_tmap0;
// SF Atom is ((32, 4), (16, 4))
const int rest_M = DIVUP(new_M, SF_TILE_ROWS);
const int rest_N = DIVUP(new_N, SF_TILE_ROWS);
auto new_M_padded = SF_TILE_ROWS * rest_M; // 128 is the number of rows in the atom.
auto new_N_padded = SF_TILE_ROWS * rest_N;
init_SF_tmap(&SFA_tmaps[workItem], SFA_ptr, new_M_padded * K / SF_BLOCK_SIZE, SF_ATOM_NUM_ELEMENTS * SF_ATOMS_PER_BLOCK_K);
init_SF_tmap(&SFB_tmaps[workItem], SFB_ptr, new_N_padded * K / SF_BLOCK_SIZE, SF_ATOM_NUM_ELEMENTS * SF_ATOMS_PER_BLOCK_K);
}
};
constexpr int PROBLEM_SIZE = 8;
CUtensorMap host_A_tmaps[PROBLEM_SIZE];
CUtensorMap host_B_tmaps[PROBLEM_SIZE];
CUtensorMap host_SFA_tmaps[PROBLEM_SIZE];
CUtensorMap host_SFB_tmaps[PROBLEM_SIZE];
half* host_C_ptrs[PROBLEM_SIZE];
int host_Ms[PROBLEM_SIZE];
int host_Ns[PROBLEM_SIZE];
populate_kernel_config(host_A_tmaps, host_B_tmaps, host_SFA_tmaps, host_SFB_tmaps, host_C_ptrs, host_Ms, host_Ns, std::integral_constant<int, PROBLEM_SIZE>{});
cudaMemcpyToSymbol(A_tmaps<PROBLEM_SIZE>, host_A_tmaps, PROBLEM_SIZE * sizeof(CUtensorMap));
cudaMemcpyToSymbol(B_tmaps<PROBLEM_SIZE>, host_B_tmaps, PROBLEM_SIZE * sizeof(CUtensorMap));
cudaMemcpyToSymbol(SFA_tmaps<PROBLEM_SIZE>, host_SFA_tmaps, PROBLEM_SIZE * sizeof(CUtensorMap));
cudaMemcpyToSymbol(SFB_tmaps<PROBLEM_SIZE>, host_SFB_tmaps, PROBLEM_SIZE * sizeof(CUtensorMap));
cudaMemcpyToSymbol(C_ptrs<PROBLEM_SIZE>, host_C_ptrs, PROBLEM_SIZE * sizeof(half*));
cudaMemcpyToSymbol(Ms<PROBLEM_SIZE>, host_Ms, PROBLEM_SIZE * sizeof(int));
cudaMemcpyToSymbol(Ns<PROBLEM_SIZE>, host_Ns, PROBLEM_SIZE * sizeof(int));
auto this_kernel = multi_gemm_kernel_p0<K, BLOCK_M, BLOCK_N, BLOCK_K, NUM_STAGES,
decltype(layout_sfa), decltype(layout_sfb),
SWAP_AB, PROBLEM_SIZE, CTA_GROUP>;
if (smem_size > 48'000)
cudaFuncSetAttribute(this_kernel, cudaFuncAttributeMaxDynamicSharedMemorySize, smem_size);
this_kernel<<<grid, tb_size, smem_size>>>();
return {outputs[0], outputs[1], outputs[2], outputs[3], outputs[4], outputs[5], outputs[6], outputs[7]};
}
template std::vector<at::Tensor> launch_gemm_cstr_p0<7168, 128, 64, 256, 9, true>(
c10::List<std::tuple<at::Tensor, at::Tensor, at::Tensor>> ABC,
c10::List<std::tuple<at::Tensor, at::Tensor>> SFAB,
c10::List<at::Tensor> outputs);
"""
CPP_SRC_0 = r"""
#include <vector>
#include <torch/script.h>
template <
int K,
int BLOCK_M,
int BLOCK_N,
int BLOCK_K,
int NUM_STAGES,
bool SWAP_AB
>
std::vector<at::Tensor> launch_gemm_cstr_p0(
c10::List<std::tuple<at::Tensor, at::Tensor, at::Tensor>> ABC,
c10::List<std::tuple<at::Tensor, at::Tensor>> SFAB,
c10::List<at::Tensor> outputs);
std::vector<at::Tensor> gemm_cstr_p0(
c10::List<std::tuple<at::Tensor, at::Tensor, at::Tensor>> ABC,
c10::List<std::tuple<at::Tensor, at::Tensor>> SFAB,
c10::List<at::Tensor> outputs) {
auto C = launch_gemm_cstr_p0<7168, 128, 64, 256, 9, true>(ABC, SFAB, outputs);
return C;
}
TORCH_LIBRARY(blockscaled_cstr_p0, m) {
m.def("gemm_cstr_p0((Tensor, Tensor, Tensor)[] ABC, (Tensor, Tensor)[] SFAB, Tensor[] outputs) -> Tensor[]");
m.impl("gemm_cstr_p0", &gemm_cstr_p0);
}
"""
CUDA_SRC_1 = r"""
using namespace cute;
constexpr int MAX_TMEM_COLS = 512;
constexpr int NUM_SMS = 148;
constexpr int NUM_BLOCKS0 = 56;
constexpr int NUM_BLOCKS1 = 112;
constexpr int NUM_BLOCKS2 = 168;
constexpr int NUM_BLOCKS3 = 112;
constexpr int NUM_BLOCKS4 = 168;
constexpr int NUM_BLOCKS5 = 168;
constexpr int NUM_BLOCKS6 = 224;
constexpr int NUM_BLOCKS7 = 168;
constexpr int GEMM_BLOCK_END0 = NUM_BLOCKS0;
constexpr int GEMM_BLOCK_END1 = GEMM_BLOCK_END0 + NUM_BLOCKS1;
constexpr int GEMM_BLOCK_END2 = GEMM_BLOCK_END1 + NUM_BLOCKS2;
constexpr int GEMM_BLOCK_END3 = GEMM_BLOCK_END2 + NUM_BLOCKS3;
constexpr int GEMM_BLOCK_END4 = GEMM_BLOCK_END3 + NUM_BLOCKS4;
constexpr int GEMM_BLOCK_END5 = GEMM_BLOCK_END4 + NUM_BLOCKS5;
constexpr int GEMM_BLOCK_END6 = GEMM_BLOCK_END5 + NUM_BLOCKS6;
constexpr int GEMM_BLOCK_END7 = GEMM_BLOCK_END6 + NUM_BLOCKS7;
constexpr int NUM_BLOCKS = GEMM_BLOCK_END7;
template<int PROBLEM_SIZE>
__constant__ CUtensorMap A_tmaps[PROBLEM_SIZE];
template<int PROBLEM_SIZE>
__constant__ CUtensorMap B_tmaps[PROBLEM_SIZE];
template<int PROBLEM_SIZE>
__constant__ CUtensorMap SFA_tmaps[PROBLEM_SIZE];
template<int PROBLEM_SIZE>
__constant__ CUtensorMap SFB_tmaps[PROBLEM_SIZE];
template<int PROBLEM_SIZE>
__constant__ half* C_ptrs[PROBLEM_SIZE];
template<int PROBLEM_SIZE>
__constant__ int Ms[PROBLEM_SIZE];
template<int PROBLEM_SIZE>
__constant__ int Ns[PROBLEM_SIZE];
template <
int K,
int BLOCK_M,
int BLOCK_N,
int BLOCK_K,
int NUM_STAGES,
class SFALayout,
class SFBLayout,
bool SWAP_AB,
int PROBLEM_SIZE,
int CTA_GROUP
>
__global__
__cluster_dims__(CTA_GROUP, 1, 1)
__launch_bounds__(BLOCK_M + 2 * WARP_SIZE) // __launch_bounds__(MAX_THREADS_PER_BLOCK, MIN_BLOCKS_PER_MP)
void multi_gemm_kernel_p1() {
// https://docs.nvidia.com/cuda/parallel-thread-execution/#tcgen05-matrix-shape
static_assert(BLOCK_M == 128); // There doesnt seem to be any code for tiling block_M
// This layout is used as TMA assumes col major, this is a colasced canonical layout.
auto asmem_layout = make_layout(make_shape(_256{}, BLOCK_M, BLOCK_K / 256), make_stride(_1{}, _256{}, 256 * BLOCK_M)); // cute::LayoutLeft{} ?
// Each member of the cluster will get a cta_rank specific offset
auto bsmem_layout = make_layout(make_shape(_256{}, BLOCK_N / CTA_GROUP, BLOCK_K / 256), make_stride(_1{}, _256{}, 256 * BLOCK_N / CTA_GROUP));
// an SF atom is ((32, 4), (16, 4)), the shmem layout is ((128 * SF_ATOMS_PER_BLOCK_K, 4)) for TMA purposes
// Since we are using CU_TENSOR_MAP_SWIZZLE_NONE, the TMA atom has shape (8, 128bits) or (8, 16) bytes, so SBO below is 8 * 16
constexpr int SF_TILE_ROWS = 128;
constexpr int SF_ATOM_NUM_ELEMENTS = 128 * 4; // cosize(SF_atom)
constexpr int SF_ATOMS_PER_BLOCK_K = BLOCK_K / 16 / 4; // each atom has 16 * 4 columns
constexpr int SF_ROWS_PER_TMEM_COL = 32; // each tmem col has 32 rows - since we are using tcgen05.cp.cta_group::1.32x128b.warpx4
using Problem = std::tuple<const CUtensorMap*, const CUtensorMap*, const CUtensorMap*, const CUtensorMap*, half*, int, int, SFALayout, SFBLayout>;
auto get_problem = [&](int problem_id) -> Problem {
auto make_sf_layout = [](int rest) {
constexpr int rest_K = K / 16 / 4;
return make_layout(make_shape(make_shape(make_shape(_32{}, _4{}), rest), make_shape(make_shape(_16{}, _4{}), rest_K)),
make_stride(make_stride(make_stride(_16{}, _4{}), rest_K * 512), make_stride(make_stride(_0{}, _1{}), _512{})));
};
const CUtensorMap* A_tmap = &A_tmaps<PROBLEM_SIZE>[problem_id];
const CUtensorMap* B_tmap = &B_tmaps<PROBLEM_SIZE>[problem_id];
const CUtensorMap* SFA_tmap = &SFA_tmaps<PROBLEM_SIZE>[problem_id];
const CUtensorMap* SFB_tmap = &SFB_tmaps<PROBLEM_SIZE>[problem_id];
half* C = C_ptrs<PROBLEM_SIZE>[problem_id];
int M = Ms<PROBLEM_SIZE>[problem_id];
int N = Ns<PROBLEM_SIZE>[problem_id];
const int rest_M = DIVUP(M, SF_TILE_ROWS);
const int rest_N = DIVUP(N, SF_TILE_ROWS);
SFALayout sfa_layout = make_sf_layout(rest_M);
SFBLayout sfb_layout = make_sf_layout(rest_N);
return {A_tmap, B_tmap, SFA_tmap, SFB_tmap, C, M, N, sfa_layout, sfb_layout};
};
// CTA rank in a cluster
int cta_rank;
asm volatile("mov.b32 %0, %%cluster_ctarank;" : "=r"(cta_rank));
const int tid = threadIdx.x;
const int lane_id = tid % WARP_SIZE;
const int warp_id = tid / WARP_SIZE;
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; // BLOCK_M * BLOCK_K elements, each element is 4 bits so divide by 2 to get bytes.
constexpr int B_size = BLOCK_N * BLOCK_K / 2 / CTA_GROUP; // each CTA only loads half of B
constexpr int SFA_size = SF_ATOM_NUM_ELEMENTS * SF_ATOMS_PER_BLOCK_K; // always copy one atom - this is ATOM_SIZE_IN_BYTES * ATOMS_PER_BLOCK_K
constexpr int SFB_size = SF_ATOM_NUM_ELEMENTS * SF_ATOMS_PER_BLOCK_K;
constexpr int STAGE_SIZE = A_size + B_size + SFA_size + SFB_size;
// set up mbarriers and tmem
// we have NUM_STAGES mbars for TMA
// NUM_STAGES mbars for MMA
// 1 mbar for mainloop
#pragma nv_diag_suppress static_var_with_dynamic_init
__shared__ int64_t mbars[NUM_STAGES * 2 + 1];
__shared__ int tmem_addr[1]; // tmem address is 32-bit
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;
const int epilogue_mbar_addr = mainloop_mbar_addr + 2 * 8;
// https://docs.nvidia.com/cuda/parallel-thread-execution/#tcgen05-mma-scale-factor-a-layout-4x
// each MMA consumes:
// - (128, 64) of A -> (128, 4) of SFA -> reshaped as (32, 4', 4) -> 4 tmem columns
constexpr int SFA_tmem = BLOCK_N * 2; // Double buffer the mma output.
constexpr int SFB_tmem = SFA_tmem + 4 * (BLOCK_K / MMA_K); // we need BLOCK_K / MMA_K many scale factors for the "tiledMMA"
if (warp_id == 0 && elect_sync()) {
for (int prefetch_idx = 0; prefetch_idx < 8; ++prefetch_idx) {
asm volatile("prefetch.tensormap [%0];" :: "l"(&A_tmaps<PROBLEM_SIZE>[prefetch_idx]) : "memory");
asm volatile("prefetch.tensormap [%0];" :: "l"(&B_tmaps<PROBLEM_SIZE>[prefetch_idx]) : "memory");
asm volatile("prefetch.tensormap [%0];" :: "l"(&SFA_tmaps<PROBLEM_SIZE>[prefetch_idx]) : "memory");
asm volatile("prefetch.tensormap [%0];" :: "l"(&SFB_tmaps<PROBLEM_SIZE>[prefetch_idx]) : "memory");
}
} else if (warp_id == 1 && elect_sync()) {
for (int i = 0; i < NUM_STAGES; i++) {
mbarrier_init(tma_mbar_addr + i * 8, CTA_GROUP); // one thread in each cluster CTA reports.
mbarrier_init(mma_mbar_addr + i * 8, 1);
}
mbarrier_init(mainloop_mbar_addr, 1); // barrier for first tmem mma buffer
mbarrier_init(mainloop_mbar_addr + 8, 1); // barrier for second tmem mma buffer
mbarrier_init(epilogue_mbar_addr, 4 * CTA_GROUP);
mbarrier_init(epilogue_mbar_addr + 8, 4 * CTA_GROUP);
asm volatile("fence.mbarrier_init.release.cluster;"); // visible to async proxy
}
if constexpr (CTA_GROUP > 1) {
// visible to all threads in a cluster
asm volatile("barrier.cluster.arrive.relaxed.aligned;");
asm volatile("barrier.cluster.wait.acquire.aligned;");
}
else {
// visible to all threads in a threadblock
__syncthreads();
}
constexpr int num_iters = K / BLOCK_K;
const int bid = blockIdx.x;
auto scheduler = [&](int global_cluster_id) -> std::tuple<int, int, int, int> {
int problem_id = 0;
int local_cluster_id = global_cluster_id;
if (global_cluster_id < GEMM_BLOCK_END0 / CTA_GROUP) {
problem_id = 0;
local_cluster_id = global_cluster_id;
} else if (global_cluster_id < GEMM_BLOCK_END1 / CTA_GROUP) {
problem_id = 1;
local_cluster_id = global_cluster_id - GEMM_BLOCK_END0 / CTA_GROUP;
} else if (global_cluster_id < GEMM_BLOCK_END2 / CTA_GROUP) {
problem_id = 2;
local_cluster_id = global_cluster_id - GEMM_BLOCK_END1 / CTA_GROUP;
} else if (global_cluster_id < GEMM_BLOCK_END3 / CTA_GROUP) {
problem_id = 3;
local_cluster_id = global_cluster_id - GEMM_BLOCK_END2 / CTA_GROUP;
} else if (global_cluster_id < GEMM_BLOCK_END4 / CTA_GROUP) {
problem_id = 4;
local_cluster_id = global_cluster_id - GEMM_BLOCK_END3 / CTA_GROUP;
} else if (global_cluster_id < GEMM_BLOCK_END5 / CTA_GROUP) {
problem_id = 5;
local_cluster_id = global_cluster_id - GEMM_BLOCK_END4 / CTA_GROUP;
} else if (global_cluster_id < GEMM_BLOCK_END6 / CTA_GROUP) {
problem_id = 6;
local_cluster_id = global_cluster_id - GEMM_BLOCK_END5 / CTA_GROUP;
} else if (global_cluster_id < GEMM_BLOCK_END7 / CTA_GROUP) {
problem_id = 7;
local_cluster_id = global_cluster_id - GEMM_BLOCK_END6 / CTA_GROUP;
}
int N = Ns<PROBLEM_SIZE>[problem_id];
const int grid_n = DIVUP(N, BLOCK_N) ;
// The cluster Ids are arranged in a (grid_m / 2 , grid_n) matrix in a row major layout.
const int bid_m = local_cluster_id / grid_n * CTA_GROUP + cta_rank;
const int bid_n = local_cluster_id % grid_n;
const int off_m = bid_m * BLOCK_M;
const int cluster_off_n = bid_n * BLOCK_N; // N offset for this cluster
const int off_n = cluster_off_n + cta_rank * (BLOCK_N / CTA_GROUP); // N offset for this CTA
return {off_m, off_n, cluster_off_n, problem_id};
};
// warp-specialization
if (warp_id == NUM_WARPS - 2 && elect_sync()) {
// TMA warp
auto issue_tma = [&](int iter_k, int stage_id, int off_m, int off_n, int problem_id) {
const int mbar_addr = (tma_mbar_addr + stage_id * 8) & 0xFEFFFFFF; // CTA0's barrier
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 problem = get_problem(problem_id);
const CUtensorMap* A_tmap = std::get<0>(problem);
const CUtensorMap* B_tmap = std::get<1>(problem);
const int M = std::get<5>(problem);
const int N = std::get<6>(problem);
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;
}
// issue TMA
const int off_k = iter_k * BLOCK_K;
tma_3d_gmem2smem<CTA_GROUP>(A_smem, A_tmap, 0, off_m, off_k / 256, mbar_addr, cache_A);
tma_3d_gmem2smem<CTA_GROUP>(B_smem, B_tmap, 0, off_n, off_k / 256, mbar_addr, cache_B);
const int rest_m = off_m / 32 / 4; // M atom is (32, 4)
const int rest_n = off_n / 32 / 4; // N atom is (32, 4)
const int rest_k = off_k / 16 / 4; // The SF atom is ((32, 4), (16, 4))
const CUtensorMap* SFA_tmap = std::get<2>(problem);
const CUtensorMap* SFB_tmap = std::get<3>(problem);
SFALayout sfa_layout = std::get<7>(problem);
SFALayout sfb_layout = std::get<8>(problem);
// Divide by 8 since underlying type is INT64
int sfa_offset = crd2idx(make_coord(make_coord(_0{}, rest_m), make_coord(_0{}, rest_k)), sfa_layout) / 8;
tma_1d_gmem2smem<CTA_GROUP>(SFA_smem, SFA_tmap, sfa_offset, mbar_addr, cache_A);
int sfb_offset = crd2idx(make_coord(make_coord(_0{}, rest_n), make_coord(_0{}, rest_k)), sfb_layout) / 8;
tma_1d_gmem2smem<CTA_GROUP>(SFB_smem, SFB_tmap, sfb_offset, mbar_addr, cache_B);
// signal TMA done .shared::cluster
asm volatile("mbarrier.arrive.expect_tx.release.cta.shared::cluster.b64 _, [%0], %1;"
:: "r"(mbar_addr), "r"(STAGE_SIZE) : "memory");
};
int mma_phase = 1; // Start with phase 1 as there are no pending MMA operations
int tma_pipeline_stage = 0;
for (int this_bid = bid; this_bid < NUM_BLOCKS; this_bid += NUM_SMS) {
auto [off_m, off_n, dummy, problem_id] = scheduler(this_bid / CTA_GROUP); // divide by CTA_GROUP to get cluster id
for (int iter_k = 0; iter_k < num_iters; iter_k++) {
// wait MMA
// NB when debugging, it will crash w/o this barrier. You cannot keep issuing TMAs w/o the
// signal the mbarrier_wait(tma_mbar_addr + stage_id * 8, tma_phase); has completed.
mbarrier_wait(mma_mbar_addr + tma_pipeline_stage * 8, mma_phase);
issue_tma(iter_k, tma_pipeline_stage, off_m, off_n, problem_id);
tma_pipeline_stage = (tma_pipeline_stage + 1) % NUM_STAGES;
if (tma_pipeline_stage == 0) {
mma_phase ^= 1;
}
}
}
}
else if (warp_id == NUM_WARPS - 1) {
// allocate tmem
const int addr = static_cast<int>(__cvta_generic_to_shared(tmem_addr));
asm volatile("tcgen05.alloc.cta_group::%2.sync.aligned.shared::cta.b32 [%0], %1;" :: "r"(addr), "r"(MAX_TMEM_COLS), "n"(CTA_GROUP));
if (cta_rank == 0 && elect_sync()) {
// MMA warp
// https://docs.nvidia.com/cuda/parallel-thread-execution/#tcgen05-instruction-descriptor
// fp4 MMA doesn't support MMA_M=64. Hence, we will use MMA_M=128 and ignore the rest.
constexpr int MMA_N = BLOCK_N;
constexpr int MMA_M = 128 * CTA_GROUP;
constexpr int16_t cta_mask = (1 << CTA_GROUP) - 1;
constexpr uint32_t i_desc = (1U << 7U) // atype=E2M1
| (1U << 10U) // btype=E2M1
| ((uint32_t)MMA_N >> 3U << 17U)
| ((uint32_t)MMA_M >> 7U << 27U)
;
int tma_phase = 0;
int mma_pipeline_stage = 0;
int epilogue_phase = 1;
int mainloop_stage = 0; // 0/1 indicates where in tmem we will put the result
for (int this_bid = bid; this_bid < NUM_BLOCKS; this_bid += NUM_SMS) {
mbarrier_wait(epilogue_mbar_addr + 8 * mainloop_stage, epilogue_phase);
auto [off_m, off_n, dummy, dummy1] = scheduler(this_bid / CTA_GROUP); // divide by CTA_GROUP to get cluster id
// Given a k slice of the scale factors which is an atom, the scale factors are stored in groups of four columns in tmem
// These will be aligned to rows 0, SF_TILE_ROWS, SF_TILE_ROWS * 2, ...
// And the sf 0..31 will be in index_col=0, 32..63 index_col =1 etc
// Given an arbitrary row_offset in M or N, in order to find the tmem column offset,
// Find the residue at the SF tile row granularity, (row_offset % SF_TILE_ROWS) and
// since each tmem col stores 32 SF, move over by the desired number of columns
// sf_tmem_offset = (row_offset % SF_TILE_ROWS) / SF_ROWS_PER_TMEM_COL;
const int M_sf_tmem_offset = (off_m % SF_TILE_ROWS) / SF_ROWS_PER_TMEM_COL;
const int N_sf_tmem_offset = (off_n % SF_TILE_ROWS) / SF_ROWS_PER_TMEM_COL;
for (int iter_k = 0; iter_k < num_iters; iter_k++) {
// wait TMA
mbarrier_wait(tma_mbar_addr + mma_pipeline_stage * 8, tma_phase);
const int A_smem = smem + mma_pipeline_stage * 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;
// 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.
auto make_desc_AB = [](int addr) -> uint64_t {
// GMMA::Layout_K_SW128_Atom<cutlass::float_e2m1_t>{} is (_8,_256):(_256,_1) so to get to the start of the next row is 256 / 2 = 128 bytes
// so to skip 8 rows, we need 8 * 128 bytes.
const int SBO = 8 * 128;
return desc_encode(addr) | (desc_encode(SBO) << 32ULL) | (1ULL << 46ULL) | (2ULL << 61ULL);
};
// no swizzling
auto make_desc_SF = [](int addr) -> uint64_t {
// Since we are using no swizzling, the atom has size (8, 16 = 128 / 8). These atoms
// are tiled (SF_ATOMS_PER_BLOCK_K = 4, 1). So the SBO is 8 * 16
const int SBO = 8 * 16;
return desc_encode(addr) | (desc_encode(SBO) << 32ULL) | (1ULL << 46ULL);
};
// tcgen05.cp -> tcgen05.mma should be pipelined correctly per PTX doc
// https://docs.nvidia.com/cuda/parallel-thread-execution/#tcgen05-memory-consistency-model-pipelined-instructions
// cutlass issues all of smem->tmem BEFORE mma
// https://github.com/NVIDIA/cutlass/blob/v4.3.2/include/cutlass/gemm/collective/sm100_blockscaled_mma_warpspecialized.hpp#L1013-L1016
constexpr uint64_t SF_desc = make_desc_SF(0);
// The below should really be SF_desc + matrix-descriptor-encode(x) , where matrix-descriptor-encode(x) = (x & 0x3FFFF) >> 4
// but (x & 0x3FFFF) >> 4 simplifies to x >> 4.
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++) {
// The below should really be SFA_desc + (uint64_t)k * matrix-descriptor-encode(512ULL); where matrix-descriptor-encode(x) = (x & 0x3FFFF) >> 4
// but (x & 0x3FFFF) >> 4 simplifies to x >> 4. So this is SFA_desc + (uint64_t)k * (512ULL >> 4ULL)
// So we are copying 512 bytes chunks , so adjust the shmem desc accordingly.
uint64_t sfa_desc = SFA_desc + (uint64_t)k * (512ULL >> 4ULL); // 4 columns, 512 bytes of 128x4 / 32x4x4
uint64_t sfb_desc = SFB_desc + (uint64_t)k * (512ULL >> 4ULL);
// tmem addresses
// Tensor Memory addresses are 32-bit wide and specify two components.
// Lane index
// Column index
// The layout is as follows:
// 31 16
// 15 0
// Lane index
// Column index
// We are using tcgen05.cp.cta_group::1.32x128b.warpx4
// See https://docs.nvidia.com/cuda/parallel-thread-execution/#tcgen05-mma-scale-factor-a-layout-4x
// so the slices ((32, 4), (16, 4)) slices of the scale factors are in the 32 rows (lanes) of tmem
tcgen05_cp_nvfp4<CTA_GROUP>(SFA_tmem + k * 4, sfa_desc);
tcgen05_cp_nvfp4<CTA_GROUP>(SFB_tmem + k * 4, sfb_desc);
}
// Compare to tiled_mma
// k1 selects the (BLOCK_M, 256) tile.
// k2 selects the (BLOCK_M, 64) tile, whose rows are swizzled.
// NOTE: this doesn't work with BLOCK_N=32, since apparently tcgen05.mma requires SFB_tmem
// to have 2-column (8-byte) alignment (looks like not documented).
// HOWEVER: works with a cluster with BLOCK_N = 64 - maybe only the issuing thread has the alignment
// requirement.
for (int k1 = 0; k1 < BLOCK_K / 256; k1++) { // BLOCK_K == 256 so this loop only executes once
for (int k2 = 0; k2 < 256 / MMA_K; k2++) { // this is the inner loop needed since MMA_K is in general < 256 = BLOCK_K
// The layout in shmem is the result of the TMA: {256, BLOCK_M, BLOCK_K / 256};
// Notice that in TMA, it’s implied that the shared memory destination has the natural contiguous layout of the given shape
// (we don’t specify shared memory stride anywhere).
// since cuTensorMapEncodeTiled assumes the tensors are col major cute::LayoutLeft
// The layout is then (256, BLOCK_M, BLOCK_K / 256) : (1, 256, 256 * BLOCK_M) elements
// https://gau-nernst.github.io/tcgen05/#decipher-tcgen05
//
// Thus given k1 and k2, the coords are (k2 * MMA_K, 0, k1)
// so the offset in elements is k1 * 256 * BLOCK_M + k2 * MMA_K = k1 * 256 * BLOCK_M + k2 * 64 (MMA_K = 64)
// so the offset in bytes = k1 * 128 * BLOCK_M + k2 * 32
// which is the expression below.
// uint64_t a_desc = make_desc_AB(A_smem + k1 * BLOCK_M * 128 + k2 * 32);
// crd2idx(make_coord(k2 * MMA_K, 0, k1), asmem_layout) is in elements, divide by 2 to get bytes
uint64_t a_desc = make_desc_AB(A_smem + crd2idx(make_coord(k2 * MMA_K, _0{}, k1), asmem_layout) / 2);
uint64_t b_desc = make_desc_AB(B_smem + crd2idx(make_coord(k2 * MMA_K, _0{}, k1), bsmem_layout) / 2);
// tmem addresses
// Tensor Memory addresses are 32-bit wide and specify two components.
// Lane index
// Column index
// The layout is as follows:
// 31 16
// 15 0
// Lane index
// Column index
int k_sf = k1 * 4 + k2; // the scale factors are in blocks of 4. This is the block idx
// k_sf is mulitplied by 4 to get us to the actual column.
const int scale_A_tmem = SFA_tmem + k_sf * 4 + M_sf_tmem_offset;
const int scale_B_tmem = SFB_tmem + k_sf * 4 + N_sf_tmem_offset;
const int enable_input_d = (k1 == 0 && k2 == 0) ? iter_k : 1;
tcgen05_mma_nvfp4<CTA_GROUP>(/*dst tmem*/ BLOCK_N * mainloop_stage, a_desc, b_desc, i_desc, scale_A_tmem, scale_B_tmem, enable_input_d);
} // for k2
} // for k1
// signal MMA done
// @@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@
// tcgen05_commit_mcast does not work if CTA_GROUP = 1
// tcgen05_commit_mcast<CTA_GROUP>(mma_mbar_addr + mma_pipeline_stage * 8, cta_mask); // signal MMA done
// @@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@
asm volatile("tcgen05.commit.cta_group::%2.mbarrier::arrive::one.shared::cluster.multicast::cluster.b64 [%0], %1;"
:: "r"(mma_mbar_addr + mma_pipeline_stage * 8), "h"(cta_mask), "n"(CTA_GROUP) : "memory");
mma_pipeline_stage = (mma_pipeline_stage + 1) % NUM_STAGES;
if (mma_pipeline_stage == 0) {
tma_phase ^= 1;
}
}
// signal mainloop done
// tcgen05_commit_mcast<CTA_GROUP>(mainloop_mbar_addr + mainloop_stage * 8, cta_mask); // signal mainloop done
asm volatile("tcgen05.commit.cta_group::%2.mbarrier::arrive::one.shared::cluster.multicast::cluster.b64 [%0], %1;"
:: "r"(mainloop_mbar_addr + mainloop_stage * 8), "h"(cta_mask), "n"(CTA_GROUP) : "memory");
mainloop_stage = (mainloop_stage + 1) % 2;
if (mainloop_stage == 0) {
epilogue_phase ^= 1;
}
} // for bid
} // cta_rank == 0 && elect_sync()
} // warp_id == NUM_WARPS - 1
else if (tid < BLOCK_M) {
// epilogue warps
int mainloop_phase = 0;
int mainloop_stage = 0; // 0/1 indicates where in tmem we will put the result
for (int this_bid = bid; this_bid < NUM_BLOCKS; this_bid += NUM_SMS) {
auto result = scheduler(this_bid / CTA_GROUP); // divide by CTA_GROUP to get cluster id
int off_m = std::get<0>(result);
int cluster_off_n = std::get<2>(result);
int problem_id = std::get<3>(result);
auto problem = get_problem(problem_id);
const int M = std::get<5>(problem);
const int N = std::get<6>(problem);
half* C_ptr = std::get<4>(problem);
// wait mainloop
mbarrier_wait(mainloop_mbar_addr + mainloop_stage * 8, mainloop_phase);
asm volatile("tcgen05.fence::after_thread_sync;");
int tmem_buffer_offset = BLOCK_N * mainloop_stage;
auto epilogue_M_major = [&]() {
// C is M-major
constexpr int WIDTH = std::min(BLOCK_N, 64); // using 128 might be slower
for (int n = 0; n < (BLOCK_N + WIDTH - 1)/ WIDTH; n++) {
float tmp[WIDTH]; // if WIDTH=128, we are using 128 registers here
if constexpr (WIDTH == 128) tcgen05_ld_32x32bx128(tmp, warp_id * 32, n * WIDTH + tmem_buffer_offset);
if constexpr (WIDTH == 64) tcgen05_ld_32x32bx64(tmp, warp_id * 32, n * WIDTH + tmem_buffer_offset);
if constexpr (WIDTH == 32) tcgen05_ld_32x32bx32(tmp, warp_id * 32, n * WIDTH + tmem_buffer_offset);
asm volatile("tcgen05.wait::ld.sync.aligned;");
for (int i = 0; i < WIDTH; i++) {
// NB row and col are transposed.
const int row = cluster_off_n + n * WIDTH + i;
const int col = off_m + tid;
if (row < N && col < M) {
C_ptr[row * M + col] = __float2half(tmp[i]);
}
}
}
};
auto epilogue_N_major = [&]() {
// C is N-major
for (int m = 0; m < 32 / 16; m++) {
float tmp[BLOCK_N / 2];
if constexpr (BLOCK_N == 128) tcgen05_ld_16x256bx16(tmp, warp_id * 32 + m * 16, tmem_buffer_offset);
if constexpr (BLOCK_N == 64) tcgen05_ld_16x256bx8(tmp, warp_id * 32 + m * 16, tmem_buffer_offset);
if constexpr (BLOCK_N == 32) tcgen05_ld_16x256bx4(tmp, warp_id * 32 + m * 16, tmem_buffer_offset);
asm volatile("tcgen05.wait::ld.sync.aligned;");
for (int i = 0; i < (BLOCK_N + 7)/ 8; i++) {
// TODO replace this with cute tensors.
const int row = off_m + warp_id * 32 + m * 16 + lane_id / 4;
const int col = cluster_off_n + i * 8 + (lane_id % 4) * 2;
if (row < M) {
reinterpret_cast<half2 *>(C_ptr + (row + 0) * N + col)[0] = __float22half2_rn({tmp[i * 4 + 0], tmp[i * 4 + 1]});
}
if (row + 8 < M) {
reinterpret_cast<half2 *>(C_ptr + (row + 8) * N + col)[0] = __float22half2_rn({tmp[i * 4 + 2], tmp[i * 4 + 3]});
}
}
}
};
if constexpr (SWAP_AB) {
epilogue_M_major();
} else {
epilogue_N_major();
}
if (elect_sync()) {
// // signal Epilogue done
// constexpr int16_t cta_mask = 3;
// tcgen05_commit_mcast<CTA_GROUP>(epilogue_mbar_addr + 8 * mainloop_stage, cta_mask); // signal Epilogue done
const int mbar_addr = (epilogue_mbar_addr + mainloop_stage * 8) & 0xFEFFFFFF;
asm volatile("mbarrier.arrive.release.cta.shared::cluster.b64 _, [%0];" :: "r"(mbar_addr) : "memory");
}
mainloop_stage = (mainloop_stage + 1) % 2;
if (mainloop_stage == 0) {
mainloop_phase ^= 1;
}
} // for bid
} //epilogue warps
// All warps.
if constexpr (CTA_GROUP > 1) {
// Dont exit w/o the peer CTA.
asm volatile("barrier.cluster.arrive.release.aligned;");
asm volatile("barrier.cluster.wait.acquire.aligned;");
} else {
__syncthreads(); // all threads finish reading data from tmem
}
// deallocate tmem. tmem address should be 0.
if (warp_id == 0) {
asm volatile("tcgen05.dealloc.cta_group::%2.sync.aligned.b32 %0, %1;" :: "r"(0), "r"(MAX_TMEM_COLS), "n"(CTA_GROUP));
}
}
template <
int K,
int BLOCK_M,
int BLOCK_N,
int BLOCK_K,
int NUM_STAGES,
bool SWAP_AB
>
std::vector<at::Tensor> launch_gemm_cstr_p1(
c10::List<std::tuple<at::Tensor, at::Tensor, at::Tensor>> ABC,
c10::List<std::tuple<at::Tensor, at::Tensor>> SFAB,
c10::List<at::Tensor> outputs) {
constexpr int CTA_GROUP = 2;
static_assert(BLOCK_K % 256 == 0);
constexpr int rest_K = K / 16 / 4;
constexpr int grid = NUM_SMS;
constexpr int tb_size = BLOCK_M + 2 * WARP_SIZE;
constexpr int AB_size = (BLOCK_M + BLOCK_N / CTA_GROUP) * (BLOCK_K / 2);
constexpr int SFAB_size = 128 * (BLOCK_K / 16) * 2;
constexpr int smem_size = (AB_size + SFAB_size) * NUM_STAGES;
// This layout is used as TMA assumes col major, this is a colasced canonical layout.
auto layout_smem_A = make_layout(make_shape(256, BLOCK_M, BLOCK_K / 256), make_stride(1, 256, 256 * BLOCK_M)); // cute::LayoutLeft{} ?
// Each member of the cluster will get a cta_rank specific offset
auto layout_smem_B = make_layout(make_shape(256, BLOCK_N / CTA_GROUP, BLOCK_K / 256), make_stride(1, 256, 256 * BLOCK_N / CTA_GROUP));
constexpr int SF_TILE_ROWS = 128;
constexpr int SF_ATOM_NUM_ELEMENTS = 128 * 4; // cosize(SF_atom)
constexpr int SF_ATOMS_PER_BLOCK_K = BLOCK_K / 16 / 4; // each atom has 16 * 4 columns
auto make_sf_layout = [&](int rest) {
return make_layout(make_shape(make_shape(make_shape(_32{}, _4{}), rest), make_shape(make_shape(_16{}, _4{}), rest_K)),
make_stride(make_stride(make_stride(_16{}, _4{}), rest_K * 512), make_stride(make_stride(_0{}, _1{}), _512{})));
};
// This is for the type only
auto layout_sfa = make_sf_layout(16);
auto layout_sfb = make_sf_layout(16);
auto populate_kernel_config = [&ABC, &SFAB, &outputs](CUtensorMap* A_tmaps, CUtensorMap* B_tmaps, CUtensorMap* SFA_tmaps, CUtensorMap* SFB_tmaps,
half** C_ptrs, int* Ms, int* Ns, auto problem_size_const) {
constexpr int PROBLEM_SIZE = decltype(problem_size_const)::value;
for (int workItem = 0; workItem < PROBLEM_SIZE; ++workItem) {
std::tuple<at::Tensor, at::Tensor, at::Tensor> abc_tuple = ABC[workItem];
auto A = std::get<0>(abc_tuple);
auto B = std::get<1>(abc_tuple);
std::tuple<at::Tensor, at::Tensor> sfab_tuple = SFAB[workItem];
auto SFA = std::get<0>(sfab_tuple);
auto SFB = std::get<1>(sfab_tuple);
at::Tensor C = outputs[workItem];
const int M = A.size(0);
const int N = B.size(0);
auto A_ptr = reinterpret_cast<const char *>(A.data_ptr());
auto B_ptr = reinterpret_cast<const char *>(B.data_ptr());
auto SFA_ptr = reinterpret_cast<const char *>(SFA.data_ptr());
auto SFB_ptr = reinterpret_cast<const char *>(SFB.data_ptr());
C_ptrs[workItem] = reinterpret_cast<half *>(C.data_ptr());
int new_M = M;
int new_N = N;
if constexpr (SWAP_AB) {
std::swap(A_ptr, B_ptr);
std::swap(SFA_ptr, SFB_ptr);
std::swap(new_M, new_N);
}
Ms[workItem] = new_M;
Ns[workItem] = new_N;
init_AB_tmap(&A_tmaps[workItem], A_ptr, new_M, K, BLOCK_M, BLOCK_K);
init_AB_tmap(&B_tmaps[workItem], B_ptr, new_N, K, BLOCK_N / CTA_GROUP, BLOCK_K);
// CUtensorMap SFA_tmap0, SFB_tmap0;
// SF Atom is ((32, 4), (16, 4))
const int rest_M = DIVUP(new_M, SF_TILE_ROWS);
const int rest_N = DIVUP(new_N, SF_TILE_ROWS);
auto new_M_padded = SF_TILE_ROWS * rest_M; // 128 is the number of rows in the atom.
auto new_N_padded = SF_TILE_ROWS * rest_N;
init_SF_tmap(&SFA_tmaps[workItem], SFA_ptr, new_M_padded * K / SF_BLOCK_SIZE, SF_ATOM_NUM_ELEMENTS * SF_ATOMS_PER_BLOCK_K);
init_SF_tmap(&SFB_tmaps[workItem], SFB_ptr, new_N_padded * K / SF_BLOCK_SIZE, SF_ATOM_NUM_ELEMENTS * SF_ATOMS_PER_BLOCK_K);
}
};
constexpr int PROBLEM_SIZE = 8;
CUtensorMap host_A_tmaps[PROBLEM_SIZE];
CUtensorMap host_B_tmaps[PROBLEM_SIZE];
CUtensorMap host_SFA_tmaps[PROBLEM_SIZE];
CUtensorMap host_SFB_tmaps[PROBLEM_SIZE];
half* host_C_ptrs[PROBLEM_SIZE];
int host_Ms[PROBLEM_SIZE];
int host_Ns[PROBLEM_SIZE];
populate_kernel_config(host_A_tmaps, host_B_tmaps, host_SFA_tmaps, host_SFB_tmaps, host_C_ptrs, host_Ms, host_Ns, std::integral_constant<int, PROBLEM_SIZE>{});
cudaMemcpyToSymbol(A_tmaps<PROBLEM_SIZE>, host_A_tmaps, PROBLEM_SIZE * sizeof(CUtensorMap));
cudaMemcpyToSymbol(B_tmaps<PROBLEM_SIZE>, host_B_tmaps, PROBLEM_SIZE * sizeof(CUtensorMap));
cudaMemcpyToSymbol(SFA_tmaps<PROBLEM_SIZE>, host_SFA_tmaps, PROBLEM_SIZE * sizeof(CUtensorMap));
cudaMemcpyToSymbol(SFB_tmaps<PROBLEM_SIZE>, host_SFB_tmaps, PROBLEM_SIZE * sizeof(CUtensorMap));
cudaMemcpyToSymbol(C_ptrs<PROBLEM_SIZE>, host_C_ptrs, PROBLEM_SIZE * sizeof(half*));
cudaMemcpyToSymbol(Ms<PROBLEM_SIZE>, host_Ms, PROBLEM_SIZE * sizeof(int));
cudaMemcpyToSymbol(Ns<PROBLEM_SIZE>, host_Ns, PROBLEM_SIZE * sizeof(int));
auto this_kernel = multi_gemm_kernel_p1<K, BLOCK_M, BLOCK_N, BLOCK_K, NUM_STAGES,
decltype(layout_sfa), decltype(layout_sfb),
SWAP_AB, PROBLEM_SIZE, CTA_GROUP>;
if (smem_size > 48'000)
cudaFuncSetAttribute(this_kernel, cudaFuncAttributeMaxDynamicSharedMemorySize, smem_size);
this_kernel<<<grid, tb_size, smem_size>>>();
return {outputs[0], outputs[1], outputs[2], outputs[3], outputs[4], outputs[5], outputs[6], outputs[7]};
}
template std::vector<at::Tensor> launch_gemm_cstr_p1<2048, 128, 64, 256, 9, true>(
c10::List<std::tuple<at::Tensor, at::Tensor, at::Tensor>> ABC,
c10::List<std::tuple<at::Tensor, at::Tensor>> SFAB,
c10::List<at::Tensor> outputs);
"""
CPP_SRC_1 = r"""
#include <vector>
#include <torch/script.h>
template <
int K,
int BLOCK_M,
int BLOCK_N,
int BLOCK_K,
int NUM_STAGES,
bool SWAP_AB
>
std::vector<at::Tensor> launch_gemm_cstr_p1(
c10::List<std::tuple<at::Tensor, at::Tensor, at::Tensor>> ABC,
c10::List<std::tuple<at::Tensor, at::Tensor>> SFAB,
c10::List<at::Tensor> outputs);
std::vector<at::Tensor> gemm_cstr_p1(
c10::List<std::tuple<at::Tensor, at::Tensor, at::Tensor>> ABC,
c10::List<std::tuple<at::Tensor, at::Tensor>> SFAB,
c10::List<at::Tensor> outputs) {
auto C = launch_gemm_cstr_p1<2048, 128, 64, 256, 9, true>(ABC, SFAB, outputs);
// #define LAUNCH(K_, BLOCK_M, BLOCK_N, BLOCK_K, NUM_STAGES) \
// else if (K == K_) C = gemm_launch<K_, BLOCK_M, BLOCK_N, BLOCK_K, NUM_STAGES>(A, B, SFA, SFB, C, buf);
// if (false) {}
// LAUNCH(16384, 128, 128, 256, 6)
// LAUNCH( 7168, 128, 64, 256, 8)
// LAUNCH( 2048, 128, 64, 256, 8)
// // the rest
// LAUNCH( 256, 128, 64, 256, 6)
// LAUNCH( 512, 128, 64, 256, 6)
// LAUNCH(1536, 128, 64, 256, 6)
// LAUNCH(2304, 128, 64, 256, 6)
// #undef LAUNCH
return C;
}
TORCH_LIBRARY(blockscaled_cstr_p1, m) {
m.def("gemm_cstr_p1((Tensor, Tensor, Tensor)[] ABC, (Tensor, Tensor)[] SFAB, Tensor[] outputs) -> Tensor[]");
m.impl("gemm_cstr_p1", &gemm_cstr_p1);
}
"""
CUDA_SRC_2 = r"""
using namespace cute;
constexpr int MAX_TMEM_COLS = 512;
constexpr int NUM_SMS = 148;
constexpr int NUM_BLOCKS0 = 72;
constexpr int NUM_BLOCKS1 = 120;
constexpr int GEMM_BLOCK_END0 = NUM_BLOCKS0;
constexpr int GEMM_BLOCK_END1 = GEMM_BLOCK_END0 + NUM_BLOCKS1;
constexpr int NUM_BLOCKS = GEMM_BLOCK_END1;
template<int PROBLEM_SIZE>
__constant__ CUtensorMap A_tmaps[PROBLEM_SIZE];
template<int PROBLEM_SIZE>
__constant__ CUtensorMap B_tmaps[PROBLEM_SIZE];
template<int PROBLEM_SIZE>
__constant__ CUtensorMap SFA_tmaps[PROBLEM_SIZE];
template<int PROBLEM_SIZE>
__constant__ CUtensorMap SFB_tmaps[PROBLEM_SIZE];
template<int PROBLEM_SIZE>
__constant__ half* C_ptrs[PROBLEM_SIZE];
template<int PROBLEM_SIZE>
__constant__ int Ms[PROBLEM_SIZE];
template<int PROBLEM_SIZE>
__constant__ int Ns[PROBLEM_SIZE];
template <
int K,
int BLOCK_M,
int BLOCK_N,
int BLOCK_K,
int NUM_STAGES,
class SFALayout,
class SFBLayout,
bool SWAP_AB,
int PROBLEM_SIZE,
int CTA_GROUP
>
__global__
__cluster_dims__(CTA_GROUP, 1, 1)
__launch_bounds__(BLOCK_M + 2 * WARP_SIZE) // __launch_bounds__(MAX_THREADS_PER_BLOCK, MIN_BLOCKS_PER_MP)
void multi_gemm_kernel_p2() {
// https://docs.nvidia.com/cuda/parallel-thread-execution/#tcgen05-matrix-shape
static_assert(BLOCK_M == 128); // There doesnt seem to be any code for tiling block_M
// This layout is used as TMA assumes col major, this is a colasced canonical layout.
auto asmem_layout = make_layout(make_shape(_256{}, BLOCK_M, BLOCK_K / 256), make_stride(_1{}, _256{}, 256 * BLOCK_M)); // cute::LayoutLeft{} ?
// Each member of the cluster will get a cta_rank specific offset
auto bsmem_layout = make_layout(make_shape(_256{}, BLOCK_N / CTA_GROUP, BLOCK_K / 256), make_stride(_1{}, _256{}, 256 * BLOCK_N / CTA_GROUP));
// an SF atom is ((32, 4), (16, 4)), the shmem layout is ((128 * SF_ATOMS_PER_BLOCK_K, 4)) for TMA purposes
// Since we are using CU_TENSOR_MAP_SWIZZLE_NONE, the TMA atom has shape (8, 128bits) or (8, 16) bytes, so SBO below is 8 * 16
constexpr int SF_TILE_ROWS = 128;
constexpr int SF_ATOM_NUM_ELEMENTS = 128 * 4; // cosize(SF_atom)
constexpr int SF_ATOMS_PER_BLOCK_K = BLOCK_K / 16 / 4; // each atom has 16 * 4 columns
constexpr int SF_ROWS_PER_TMEM_COL = 32; // each tmem col has 32 rows - since we are using tcgen05.cp.cta_group::1.32x128b.warpx4
using Problem = std::tuple<const CUtensorMap*, const CUtensorMap*, const CUtensorMap*, const CUtensorMap*, half*, int, int, SFALayout, SFBLayout>;
auto get_problem = [&](int problem_id) -> Problem {
auto make_sf_layout = [](int rest) {
constexpr int rest_K = K / 16 / 4;
return make_layout(make_shape(make_shape(make_shape(_32{}, _4{}), rest), make_shape(make_shape(_16{}, _4{}), rest_K)),
make_stride(make_stride(make_stride(_16{}, _4{}), rest_K * 512), make_stride(make_stride(_0{}, _1{}), _512{})));
};
const CUtensorMap* A_tmap = &A_tmaps<PROBLEM_SIZE>[problem_id];
const CUtensorMap* B_tmap = &B_tmaps<PROBLEM_SIZE>[problem_id];
const CUtensorMap* SFA_tmap = &SFA_tmaps<PROBLEM_SIZE>[problem_id];
const CUtensorMap* SFB_tmap = &SFB_tmaps<PROBLEM_SIZE>[problem_id];
half* C = C_ptrs<PROBLEM_SIZE>[problem_id];
int M = Ms<PROBLEM_SIZE>[problem_id];
int N = Ns<PROBLEM_SIZE>[problem_id];
const int rest_M = DIVUP(M, SF_TILE_ROWS);
const int rest_N = DIVUP(N, SF_TILE_ROWS);
SFALayout sfa_layout = make_sf_layout(rest_M);
SFBLayout sfb_layout = make_sf_layout(rest_N);
return {A_tmap, B_tmap, SFA_tmap, SFB_tmap, C, M, N, sfa_layout, sfb_layout};
};
// CTA rank in a cluster
int cta_rank;
asm volatile("mov.b32 %0, %%cluster_ctarank;" : "=r"(cta_rank));
const int tid = threadIdx.x;
const int lane_id = tid % WARP_SIZE;
const int warp_id = tid / WARP_SIZE;
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; // BLOCK_M * BLOCK_K elements, each element is 4 bits so divide by 2 to get bytes.
constexpr int B_size = BLOCK_N * BLOCK_K / 2 / CTA_GROUP; // each CTA only loads half of B
constexpr int SFA_size = SF_ATOM_NUM_ELEMENTS * SF_ATOMS_PER_BLOCK_K; // always copy one atom - this is ATOM_SIZE_IN_BYTES * ATOMS_PER_BLOCK_K
constexpr int SFB_size = SF_ATOM_NUM_ELEMENTS * SF_ATOMS_PER_BLOCK_K;
constexpr int STAGE_SIZE = A_size + B_size + SFA_size + SFB_size;
// set up mbarriers and tmem
// we have NUM_STAGES mbars for TMA
// NUM_STAGES mbars for MMA
// 1 mbar for mainloop
#pragma nv_diag_suppress static_var_with_dynamic_init
__shared__ int64_t mbars[NUM_STAGES * 2 + 1];
__shared__ int tmem_addr[1]; // tmem address is 32-bit
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;
const int epilogue_mbar_addr = mainloop_mbar_addr + 2 * 8;
// https://docs.nvidia.com/cuda/parallel-thread-execution/#tcgen05-mma-scale-factor-a-layout-4x
// each MMA consumes:
// - (128, 64) of A -> (128, 4) of SFA -> reshaped as (32, 4', 4) -> 4 tmem columns
constexpr int SFA_tmem = BLOCK_N * 2; // Double buffer the mma output.
constexpr int SFB_tmem = SFA_tmem + 4 * (BLOCK_K / MMA_K); // we need BLOCK_K / MMA_K many scale factors for the "tiledMMA"
if (warp_id == 0 && elect_sync()) {
for (int prefetch_idx = 0; prefetch_idx < 8; ++prefetch_idx) {
asm volatile("prefetch.tensormap [%0];" :: "l"(&A_tmaps<PROBLEM_SIZE>[prefetch_idx]) : "memory");
asm volatile("prefetch.tensormap [%0];" :: "l"(&B_tmaps<PROBLEM_SIZE>[prefetch_idx]) : "memory");
asm volatile("prefetch.tensormap [%0];" :: "l"(&SFA_tmaps<PROBLEM_SIZE>[prefetch_idx]) : "memory");
asm volatile("prefetch.tensormap [%0];" :: "l"(&SFB_tmaps<PROBLEM_SIZE>[prefetch_idx]) : "memory");
}
} else if (warp_id == 1 && elect_sync()) {
for (int i = 0; i < NUM_STAGES; i++) {
mbarrier_init(tma_mbar_addr + i * 8, CTA_GROUP); // one thread in each cluster CTA reports.
mbarrier_init(mma_mbar_addr + i * 8, 1);
}
mbarrier_init(mainloop_mbar_addr, 1); // barrier for first tmem mma buffer
mbarrier_init(mainloop_mbar_addr + 8, 1); // barrier for second tmem mma buffer
mbarrier_init(epilogue_mbar_addr, 4 * CTA_GROUP);
mbarrier_init(epilogue_mbar_addr + 8, 4 * CTA_GROUP);
asm volatile("fence.mbarrier_init.release.cluster;"); // visible to async proxy
}
if constexpr (CTA_GROUP > 1) {
// visible to all threads in a cluster
asm volatile("barrier.cluster.arrive.relaxed.aligned;");
asm volatile("barrier.cluster.wait.acquire.aligned;");
}
else {
// visible to all threads in a threadblock
__syncthreads();
}
constexpr int num_iters = K / BLOCK_K;
const int bid = blockIdx.x;
auto scheduler = [&](int global_cluster_id) -> std::tuple<int, int, int, int> {
int problem_id = 0;
int local_cluster_id = global_cluster_id;
if (global_cluster_id < GEMM_BLOCK_END0 / CTA_GROUP) {
problem_id = 0;
local_cluster_id = global_cluster_id;
} else if (global_cluster_id < GEMM_BLOCK_END1 / CTA_GROUP) {
problem_id = 1;
local_cluster_id = global_cluster_id - GEMM_BLOCK_END0 / CTA_GROUP;
}
int N = Ns<PROBLEM_SIZE>[problem_id];
const int grid_n = DIVUP(N, BLOCK_N) ;
// The cluster Ids are arranged in a (grid_m / 2 , grid_n) matrix in a row major layout.
const int bid_m = local_cluster_id / grid_n * CTA_GROUP + cta_rank;
const int bid_n = local_cluster_id % grid_n;
const int off_m = bid_m * BLOCK_M;
const int cluster_off_n = bid_n * BLOCK_N; // N offset for this cluster
const int off_n = cluster_off_n + cta_rank * (BLOCK_N / CTA_GROUP); // N offset for this CTA
return {off_m, off_n, cluster_off_n, problem_id};
};
// warp-specialization
if (warp_id == NUM_WARPS - 2 && elect_sync()) {
// TMA warp
auto issue_tma = [&](int iter_k, int stage_id, int off_m, int off_n, int problem_id) {
const int mbar_addr = (tma_mbar_addr + stage_id * 8) & 0xFEFFFFFF; // CTA0's barrier
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 problem = get_problem(problem_id);
const CUtensorMap* A_tmap = std::get<0>(problem);
const CUtensorMap* B_tmap = std::get<1>(problem);
const int M = std::get<5>(problem);
const int N = std::get<6>(problem);
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;
}
// issue TMA
const int off_k = iter_k * BLOCK_K;
tma_3d_gmem2smem<CTA_GROUP>(A_smem, A_tmap, 0, off_m, off_k / 256, mbar_addr, cache_A);
tma_3d_gmem2smem<CTA_GROUP>(B_smem, B_tmap, 0, off_n, off_k / 256, mbar_addr, cache_B);
const int rest_m = off_m / 32 / 4; // M atom is (32, 4)
const int rest_n = off_n / 32 / 4; // N atom is (32, 4)
const int rest_k = off_k / 16 / 4; // The SF atom is ((32, 4), (16, 4))
const CUtensorMap* SFA_tmap = std::get<2>(problem);
const CUtensorMap* SFB_tmap = std::get<3>(problem);
SFALayout sfa_layout = std::get<7>(problem);
SFALayout sfb_layout = std::get<8>(problem);
// Divide by 8 since underlying type is INT64
int sfa_offset = crd2idx(make_coord(make_coord(_0{}, rest_m), make_coord(_0{}, rest_k)), sfa_layout) / 8;
tma_1d_gmem2smem<CTA_GROUP>(SFA_smem, SFA_tmap, sfa_offset, mbar_addr, cache_A);
int sfb_offset = crd2idx(make_coord(make_coord(_0{}, rest_n), make_coord(_0{}, rest_k)), sfb_layout) / 8;
tma_1d_gmem2smem<CTA_GROUP>(SFB_smem, SFB_tmap, sfb_offset, mbar_addr, cache_B);
// signal TMA done .shared::cluster
asm volatile("mbarrier.arrive.expect_tx.release.cta.shared::cluster.b64 _, [%0], %1;"
:: "r"(mbar_addr), "r"(STAGE_SIZE) : "memory");
};
int mma_phase = 1; // Start with phase 1 as there are no pending MMA operations
int tma_pipeline_stage = 0;
for (int this_bid = bid; this_bid < NUM_BLOCKS; this_bid += NUM_SMS) {
auto [off_m, off_n, dummy, problem_id] = scheduler(this_bid / CTA_GROUP); // divide by CTA_GROUP to get cluster id
for (int iter_k = 0; iter_k < num_iters; iter_k++) {
// wait MMA
// NB when debugging, it will crash w/o this barrier. You cannot keep issuing TMAs w/o the
// signal the mbarrier_wait(tma_mbar_addr + stage_id * 8, tma_phase); has completed.
mbarrier_wait(mma_mbar_addr + tma_pipeline_stage * 8, mma_phase);
issue_tma(iter_k, tma_pipeline_stage, off_m, off_n, problem_id);
tma_pipeline_stage = (tma_pipeline_stage + 1) % NUM_STAGES;
if (tma_pipeline_stage == 0) {
mma_phase ^= 1;
}
}
}
}
else if (warp_id == NUM_WARPS - 1) {
// allocate tmem
const int addr = static_cast<int>(__cvta_generic_to_shared(tmem_addr));
asm volatile("tcgen05.alloc.cta_group::%2.sync.aligned.shared::cta.b32 [%0], %1;" :: "r"(addr), "r"(MAX_TMEM_COLS), "n"(CTA_GROUP));
if (cta_rank == 0 && elect_sync()) {
// MMA warp
// https://docs.nvidia.com/cuda/parallel-thread-execution/#tcgen05-instruction-descriptor
// fp4 MMA doesn't support MMA_M=64. Hence, we will use MMA_M=128 and ignore the rest.
constexpr int MMA_N = BLOCK_N;
constexpr int MMA_M = 128 * CTA_GROUP;
constexpr int16_t cta_mask = (1 << CTA_GROUP) - 1;
constexpr uint32_t i_desc = (1U << 7U) // atype=E2M1
| (1U << 10U) // btype=E2M1
| ((uint32_t)MMA_N >> 3U << 17U)
| ((uint32_t)MMA_M >> 7U << 27U)
;
int tma_phase = 0;
int mma_pipeline_stage = 0;
int epilogue_phase = 1;
int mainloop_stage = 0; // 0/1 indicates where in tmem we will put the result
for (int this_bid = bid; this_bid < NUM_BLOCKS; this_bid += NUM_SMS) {
mbarrier_wait(epilogue_mbar_addr + 8 * mainloop_stage, epilogue_phase);
auto [off_m, off_n, dummy, dummy1] = scheduler(this_bid / CTA_GROUP); // divide by CTA_GROUP to get cluster id
// Given a k slice of the scale factors which is an atom, the scale factors are stored in groups of four columns in tmem
// These will be aligned to rows 0, SF_TILE_ROWS, SF_TILE_ROWS * 2, ...
// And the sf 0..31 will be in index_col=0, 32..63 index_col =1 etc
// Given an arbitrary row_offset in M or N, in order to find the tmem column offset,
// Find the residue at the SF tile row granularity, (row_offset % SF_TILE_ROWS) and
// since each tmem col stores 32 SF, move over by the desired number of columns
// sf_tmem_offset = (row_offset % SF_TILE_ROWS) / SF_ROWS_PER_TMEM_COL;
const int M_sf_tmem_offset = (off_m % SF_TILE_ROWS) / SF_ROWS_PER_TMEM_COL;
const int N_sf_tmem_offset = (off_n % SF_TILE_ROWS) / SF_ROWS_PER_TMEM_COL;
for (int iter_k = 0; iter_k < num_iters; iter_k++) {
// wait TMA
mbarrier_wait(tma_mbar_addr + mma_pipeline_stage * 8, tma_phase);
const int A_smem = smem + mma_pipeline_stage * 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;
// 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.
auto make_desc_AB = [](int addr) -> uint64_t {
// GMMA::Layout_K_SW128_Atom<cutlass::float_e2m1_t>{} is (_8,_256):(_256,_1) so to get to the start of the next row is 256 / 2 = 128 bytes
// so to skip 8 rows, we need 8 * 128 bytes.
const int SBO = 8 * 128;
return desc_encode(addr) | (desc_encode(SBO) << 32ULL) | (1ULL << 46ULL) | (2ULL << 61ULL);
};
// no swizzling
auto make_desc_SF = [](int addr) -> uint64_t {
// Since we are using no swizzling, the atom has size (8, 16 = 128 / 8). These atoms
// are tiled (SF_ATOMS_PER_BLOCK_K = 4, 1). So the SBO is 8 * 16
const int SBO = 8 * 16;
return desc_encode(addr) | (desc_encode(SBO) << 32ULL) | (1ULL << 46ULL);
};
// tcgen05.cp -> tcgen05.mma should be pipelined correctly per PTX doc
// https://docs.nvidia.com/cuda/parallel-thread-execution/#tcgen05-memory-consistency-model-pipelined-instructions
// cutlass issues all of smem->tmem BEFORE mma
// https://github.com/NVIDIA/cutlass/blob/v4.3.2/include/cutlass/gemm/collective/sm100_blockscaled_mma_warpspecialized.hpp#L1013-L1016
constexpr uint64_t SF_desc = make_desc_SF(0);
// The below should really be SF_desc + matrix-descriptor-encode(x) , where matrix-descriptor-encode(x) = (x & 0x3FFFF) >> 4
// but (x & 0x3FFFF) >> 4 simplifies to x >> 4.
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++) {
// The below should really be SFA_desc + (uint64_t)k * matrix-descriptor-encode(512ULL); where matrix-descriptor-encode(x) = (x & 0x3FFFF) >> 4
// but (x & 0x3FFFF) >> 4 simplifies to x >> 4. So this is SFA_desc + (uint64_t)k * (512ULL >> 4ULL)
// So we are copying 512 bytes chunks , so adjust the shmem desc accordingly.
uint64_t sfa_desc = SFA_desc + (uint64_t)k * (512ULL >> 4ULL); // 4 columns, 512 bytes of 128x4 / 32x4x4
uint64_t sfb_desc = SFB_desc + (uint64_t)k * (512ULL >> 4ULL);
// tmem addresses
// Tensor Memory addresses are 32-bit wide and specify two components.
// Lane index
// Column index
// The layout is as follows:
// 31 16
// 15 0
// Lane index
// Column index
// We are using tcgen05.cp.cta_group::1.32x128b.warpx4
// See https://docs.nvidia.com/cuda/parallel-thread-execution/#tcgen05-mma-scale-factor-a-layout-4x
// so the slices ((32, 4), (16, 4)) slices of the scale factors are in the 32 rows (lanes) of tmem
tcgen05_cp_nvfp4<CTA_GROUP>(SFA_tmem + k * 4, sfa_desc);
tcgen05_cp_nvfp4<CTA_GROUP>(SFB_tmem + k * 4, sfb_desc);
}
// Compare to tiled_mma
// k1 selects the (BLOCK_M, 256) tile.
// k2 selects the (BLOCK_M, 64) tile, whose rows are swizzled.
// NOTE: this doesn't work with BLOCK_N=32, since apparently tcgen05.mma requires SFB_tmem
// to have 2-column (8-byte) alignment (looks like not documented).
// HOWEVER: works with a cluster with BLOCK_N = 64 - maybe only the issuing thread has the alignment
// requirement.
for (int k1 = 0; k1 < BLOCK_K / 256; k1++) { // BLOCK_K == 256 so this loop only executes once
for (int k2 = 0; k2 < 256 / MMA_K; k2++) { // this is the inner loop needed since MMA_K is in general < 256 = BLOCK_K
// The layout in shmem is the result of the TMA: {256, BLOCK_M, BLOCK_K / 256};
// Notice that in TMA, it’s implied that the shared memory destination has the natural contiguous layout of the given shape
// (we don’t specify shared memory stride anywhere).
// since cuTensorMapEncodeTiled assumes the tensors are col major cute::LayoutLeft
// The layout is then (256, BLOCK_M, BLOCK_K / 256) : (1, 256, 256 * BLOCK_M) elements
// https://gau-nernst.github.io/tcgen05/#decipher-tcgen05
//
// Thus given k1 and k2, the coords are (k2 * MMA_K, 0, k1)
// so the offset in elements is k1 * 256 * BLOCK_M + k2 * MMA_K = k1 * 256 * BLOCK_M + k2 * 64 (MMA_K = 64)
// so the offset in bytes = k1 * 128 * BLOCK_M + k2 * 32
// which is the expression below.
// uint64_t a_desc = make_desc_AB(A_smem + k1 * BLOCK_M * 128 + k2 * 32);
// crd2idx(make_coord(k2 * MMA_K, 0, k1), asmem_layout) is in elements, divide by 2 to get bytes
uint64_t a_desc = make_desc_AB(A_smem + crd2idx(make_coord(k2 * MMA_K, _0{}, k1), asmem_layout) / 2);
uint64_t b_desc = make_desc_AB(B_smem + crd2idx(make_coord(k2 * MMA_K, _0{}, k1), bsmem_layout) / 2);
// tmem addresses
// Tensor Memory addresses are 32-bit wide and specify two components.
// Lane index
// Column index
// The layout is as follows:
// 31 16
// 15 0
// Lane index
// Column index
int k_sf = k1 * 4 + k2; // the scale factors are in blocks of 4. This is the block idx
// k_sf is mulitplied by 4 to get us to the actual column.
const int scale_A_tmem = SFA_tmem + k_sf * 4 + M_sf_tmem_offset;
const int scale_B_tmem = SFB_tmem + k_sf * 4 + N_sf_tmem_offset;
const int enable_input_d = (k1 == 0 && k2 == 0) ? iter_k : 1;
tcgen05_mma_nvfp4<CTA_GROUP>(/*dst tmem*/ BLOCK_N * mainloop_stage, a_desc, b_desc, i_desc, scale_A_tmem, scale_B_tmem, enable_input_d);
} // for k2
} // for k1
// signal MMA done
// @@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@
// tcgen05_commit_mcast does not work if CTA_GROUP = 1
// tcgen05_commit_mcast<CTA_GROUP>(mma_mbar_addr + mma_pipeline_stage * 8, cta_mask); // signal MMA done
// @@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@
asm volatile("tcgen05.commit.cta_group::%2.mbarrier::arrive::one.shared::cluster.multicast::cluster.b64 [%0], %1;"
:: "r"(mma_mbar_addr + mma_pipeline_stage * 8), "h"(cta_mask), "n"(CTA_GROUP) : "memory");
mma_pipeline_stage = (mma_pipeline_stage + 1) % NUM_STAGES;
if (mma_pipeline_stage == 0) {
tma_phase ^= 1;
}
}
// signal mainloop done
// tcgen05_commit_mcast<CTA_GROUP>(mainloop_mbar_addr + mainloop_stage * 8, cta_mask); // signal mainloop done
asm volatile("tcgen05.commit.cta_group::%2.mbarrier::arrive::one.shared::cluster.multicast::cluster.b64 [%0], %1;"
:: "r"(mainloop_mbar_addr + mainloop_stage * 8), "h"(cta_mask), "n"(CTA_GROUP) : "memory");
mainloop_stage = (mainloop_stage + 1) % 2;
if (mainloop_stage == 0) {
epilogue_phase ^= 1;
}
} // for bid
} // cta_rank == 0 && elect_sync()
} // warp_id == NUM_WARPS - 1
else if (tid < BLOCK_M) {
// epilogue warps
int mainloop_phase = 0;
int mainloop_stage = 0; // 0/1 indicates where in tmem we will put the result
for (int this_bid = bid; this_bid < NUM_BLOCKS; this_bid += NUM_SMS) {
auto result = scheduler(this_bid / CTA_GROUP); // divide by CTA_GROUP to get cluster id
int off_m = std::get<0>(result);
int cluster_off_n = std::get<2>(result);
int problem_id = std::get<3>(result);
auto problem = get_problem(problem_id);
const int M = std::get<5>(problem);
const int N = std::get<6>(problem);
half* C_ptr = std::get<4>(problem);
// wait mainloop
mbarrier_wait(mainloop_mbar_addr + mainloop_stage * 8, mainloop_phase);
asm volatile("tcgen05.fence::after_thread_sync;");
int tmem_buffer_offset = BLOCK_N * mainloop_stage;
auto epilogue_M_major = [&]() {
// C is M-major
constexpr int WIDTH = std::min(BLOCK_N, 64); // using 128 might be slower
for (int n = 0; n < (BLOCK_N + WIDTH - 1)/ WIDTH; n++) {
float tmp[WIDTH]; // if WIDTH=128, we are using 128 registers here
if constexpr (WIDTH == 128) tcgen05_ld_32x32bx128(tmp, warp_id * 32, n * WIDTH + tmem_buffer_offset);
if constexpr (WIDTH == 64) tcgen05_ld_32x32bx64(tmp, warp_id * 32, n * WIDTH + tmem_buffer_offset);
if constexpr (WIDTH == 32) tcgen05_ld_32x32bx32(tmp, warp_id * 32, n * WIDTH + tmem_buffer_offset);
asm volatile("tcgen05.wait::ld.sync.aligned;");
for (int i = 0; i < WIDTH; i++) {
// NB row and col are transposed.
const int row = cluster_off_n + n * WIDTH + i;
const int col = off_m + tid;
if (row < N && col < M) {
C_ptr[row * M + col] = __float2half(tmp[i]);
}
}
}
};
auto epilogue_N_major = [&]() {
// C is N-major
for (int m = 0; m < 32 / 16; m++) {
float tmp[BLOCK_N / 2];
if constexpr (BLOCK_N == 128) tcgen05_ld_16x256bx16(tmp, warp_id * 32 + m * 16, tmem_buffer_offset);
if constexpr (BLOCK_N == 64) tcgen05_ld_16x256bx8(tmp, warp_id * 32 + m * 16, tmem_buffer_offset);
if constexpr (BLOCK_N == 32) tcgen05_ld_16x256bx4(tmp, warp_id * 32 + m * 16, tmem_buffer_offset);
asm volatile("tcgen05.wait::ld.sync.aligned;");
for (int i = 0; i < (BLOCK_N + 7)/ 8; i++) {
// TODO replace this with cute tensors.
const int row = off_m + warp_id * 32 + m * 16 + lane_id / 4;
const int col = cluster_off_n + i * 8 + (lane_id % 4) * 2;
if (row < M) {
reinterpret_cast<half2 *>(C_ptr + (row + 0) * N + col)[0] = __float22half2_rn({tmp[i * 4 + 0], tmp[i * 4 + 1]});
}
if (row + 8 < M) {
reinterpret_cast<half2 *>(C_ptr + (row + 8) * N + col)[0] = __float22half2_rn({tmp[i * 4 + 2], tmp[i * 4 + 3]});
}
}
}
};
if constexpr (SWAP_AB) {
epilogue_M_major();
} else {
epilogue_N_major();
}
if (elect_sync()) {
// // signal Epilogue done
// constexpr int16_t cta_mask = 3;
// tcgen05_commit_mcast<CTA_GROUP>(epilogue_mbar_addr + 8 * mainloop_stage, cta_mask); // signal Epilogue done
const int mbar_addr = (epilogue_mbar_addr + mainloop_stage * 8) & 0xFEFFFFFF;
asm volatile("mbarrier.arrive.release.cta.shared::cluster.b64 _, [%0];" :: "r"(mbar_addr) : "memory");
}
mainloop_stage = (mainloop_stage + 1) % 2;
if (mainloop_stage == 0) {
mainloop_phase ^= 1;
}
} // for bid
} //epilogue warps
// All warps.
if constexpr (CTA_GROUP > 1) {
// Dont exit w/o the peer CTA.
asm volatile("barrier.cluster.arrive.release.aligned;");
asm volatile("barrier.cluster.wait.acquire.aligned;");
} else {
__syncthreads(); // all threads finish reading data from tmem
}
// deallocate tmem. tmem address should be 0.
if (warp_id == 0) {
asm volatile("tcgen05.dealloc.cta_group::%2.sync.aligned.b32 %0, %1;" :: "r"(0), "r"(MAX_TMEM_COLS), "n"(CTA_GROUP));
}
}
template <
int K,
int BLOCK_M,
int BLOCK_N,
int BLOCK_K,
int NUM_STAGES,
bool SWAP_AB
>
std::vector<at::Tensor> launch_gemm_cstr_p2(
c10::List<std::tuple<at::Tensor, at::Tensor, at::Tensor>> ABC,
c10::List<std::tuple<at::Tensor, at::Tensor>> SFAB,
c10::List<at::Tensor> outputs) {
constexpr int CTA_GROUP = 2;
static_assert(BLOCK_K % 256 == 0);
constexpr int rest_K = K / 16 / 4;
constexpr int grid = NUM_SMS;
constexpr int tb_size = BLOCK_M + 2 * WARP_SIZE;
constexpr int AB_size = (BLOCK_M + BLOCK_N / CTA_GROUP) * (BLOCK_K / 2);
constexpr int SFAB_size = 128 * (BLOCK_K / 16) * 2;
constexpr int smem_size = (AB_size + SFAB_size) * NUM_STAGES;
// This layout is used as TMA assumes col major, this is a colasced canonical layout.
auto layout_smem_A = make_layout(make_shape(256, BLOCK_M, BLOCK_K / 256), make_stride(1, 256, 256 * BLOCK_M)); // cute::LayoutLeft{} ?
// Each member of the cluster will get a cta_rank specific offset
auto layout_smem_B = make_layout(make_shape(256, BLOCK_N / CTA_GROUP, BLOCK_K / 256), make_stride(1, 256, 256 * BLOCK_N / CTA_GROUP));
constexpr int SF_TILE_ROWS = 128;
constexpr int SF_ATOM_NUM_ELEMENTS = 128 * 4; // cosize(SF_atom)
constexpr int SF_ATOMS_PER_BLOCK_K = BLOCK_K / 16 / 4; // each atom has 16 * 4 columns
auto make_sf_layout = [&](int rest) {
return make_layout(make_shape(make_shape(make_shape(_32{}, _4{}), rest), make_shape(make_shape(_16{}, _4{}), rest_K)),
make_stride(make_stride(make_stride(_16{}, _4{}), rest_K * 512), make_stride(make_stride(_0{}, _1{}), _512{})));
};
// This is for the type only
auto layout_sfa = make_sf_layout(16);
auto layout_sfb = make_sf_layout(16);
auto populate_kernel_config = [&ABC, &SFAB, &outputs](CUtensorMap* A_tmaps, CUtensorMap* B_tmaps, CUtensorMap* SFA_tmaps, CUtensorMap* SFB_tmaps,
half** C_ptrs, int* Ms, int* Ns, auto problem_size_const) {
constexpr int PROBLEM_SIZE = decltype(problem_size_const)::value;
for (int workItem = 0; workItem < PROBLEM_SIZE; ++workItem) {
std::tuple<at::Tensor, at::Tensor, at::Tensor> abc_tuple = ABC[workItem];
auto A = std::get<0>(abc_tuple);
auto B = std::get<1>(abc_tuple);
std::tuple<at::Tensor, at::Tensor> sfab_tuple = SFAB[workItem];
auto SFA = std::get<0>(sfab_tuple);
auto SFB = std::get<1>(sfab_tuple);
at::Tensor C = outputs[workItem];
const int M = A.size(0);
const int N = B.size(0);
auto A_ptr = reinterpret_cast<const char *>(A.data_ptr());
auto B_ptr = reinterpret_cast<const char *>(B.data_ptr());
auto SFA_ptr = reinterpret_cast<const char *>(SFA.data_ptr());
auto SFB_ptr = reinterpret_cast<const char *>(SFB.data_ptr());
C_ptrs[workItem] = reinterpret_cast<half *>(C.data_ptr());
int new_M = M;
int new_N = N;
if constexpr (SWAP_AB) {
std::swap(A_ptr, B_ptr);
std::swap(SFA_ptr, SFB_ptr);
std::swap(new_M, new_N);
}
Ms[workItem] = new_M;
Ns[workItem] = new_N;
init_AB_tmap(&A_tmaps[workItem], A_ptr, new_M, K, BLOCK_M, BLOCK_K);
init_AB_tmap(&B_tmaps[workItem], B_ptr, new_N, K, BLOCK_N / CTA_GROUP, BLOCK_K);
// CUtensorMap SFA_tmap0, SFB_tmap0;
// SF Atom is ((32, 4), (16, 4))
const int rest_M = DIVUP(new_M, SF_TILE_ROWS);
const int rest_N = DIVUP(new_N, SF_TILE_ROWS);
auto new_M_padded = SF_TILE_ROWS * rest_M; // 128 is the number of rows in the atom.
auto new_N_padded = SF_TILE_ROWS * rest_N;
init_SF_tmap(&SFA_tmaps[workItem], SFA_ptr, new_M_padded * K / SF_BLOCK_SIZE, SF_ATOM_NUM_ELEMENTS * SF_ATOMS_PER_BLOCK_K);
init_SF_tmap(&SFB_tmaps[workItem], SFB_ptr, new_N_padded * K / SF_BLOCK_SIZE, SF_ATOM_NUM_ELEMENTS * SF_ATOMS_PER_BLOCK_K);
}
};
constexpr int PROBLEM_SIZE = 2;
CUtensorMap host_A_tmaps[PROBLEM_SIZE];
CUtensorMap host_B_tmaps[PROBLEM_SIZE];
CUtensorMap host_SFA_tmaps[PROBLEM_SIZE];
CUtensorMap host_SFB_tmaps[PROBLEM_SIZE];
half* host_C_ptrs[PROBLEM_SIZE];
int host_Ms[PROBLEM_SIZE];
int host_Ns[PROBLEM_SIZE];
populate_kernel_config(host_A_tmaps, host_B_tmaps, host_SFA_tmaps, host_SFB_tmaps, host_C_ptrs, host_Ms, host_Ns, std::integral_constant<int, PROBLEM_SIZE>{});
cudaMemcpyToSymbol(A_tmaps<PROBLEM_SIZE>, host_A_tmaps, PROBLEM_SIZE * sizeof(CUtensorMap));
cudaMemcpyToSymbol(B_tmaps<PROBLEM_SIZE>, host_B_tmaps, PROBLEM_SIZE * sizeof(CUtensorMap));
cudaMemcpyToSymbol(SFA_tmaps<PROBLEM_SIZE>, host_SFA_tmaps, PROBLEM_SIZE * sizeof(CUtensorMap));
cudaMemcpyToSymbol(SFB_tmaps<PROBLEM_SIZE>, host_SFB_tmaps, PROBLEM_SIZE * sizeof(CUtensorMap));
cudaMemcpyToSymbol(C_ptrs<PROBLEM_SIZE>, host_C_ptrs, PROBLEM_SIZE * sizeof(half*));
cudaMemcpyToSymbol(Ms<PROBLEM_SIZE>, host_Ms, PROBLEM_SIZE * sizeof(int));
cudaMemcpyToSymbol(Ns<PROBLEM_SIZE>, host_Ns, PROBLEM_SIZE * sizeof(int));
auto this_kernel = multi_gemm_kernel_p2<K, BLOCK_M, BLOCK_N, BLOCK_K, NUM_STAGES,
decltype(layout_sfa), decltype(layout_sfb),
SWAP_AB, PROBLEM_SIZE, CTA_GROUP>;
if (smem_size > 48'000)
cudaFuncSetAttribute(this_kernel, cudaFuncAttributeMaxDynamicSharedMemorySize, smem_size);
this_kernel<<<grid, tb_size, smem_size>>>();
return {outputs[0], outputs[1]};
}
template std::vector<at::Tensor> launch_gemm_cstr_p2<4096, 128, 64, 256, 9, true>(
c10::List<std::tuple<at::Tensor, at::Tensor, at::Tensor>> ABC,
c10::List<std::tuple<at::Tensor, at::Tensor>> SFAB,
c10::List<at::Tensor> outputs);
"""
CPP_SRC_2 = r"""
#include <vector>
#include <torch/script.h>
template <
int K,
int BLOCK_M,
int BLOCK_N,
int BLOCK_K,
int NUM_STAGES,
bool SWAP_AB
>
std::vector<at::Tensor> launch_gemm_cstr_p2(
c10::List<std::tuple<at::Tensor, at::Tensor, at::Tensor>> ABC,
c10::List<std::tuple<at::Tensor, at::Tensor>> SFAB,
c10::List<at::Tensor> outputs);
std::vector<at::Tensor> gemm_cstr_p2(
c10::List<std::tuple<at::Tensor, at::Tensor, at::Tensor>> ABC,
c10::List<std::tuple<at::Tensor, at::Tensor>> SFAB,
c10::List<at::Tensor> outputs) {
auto C = launch_gemm_cstr_p2<4096, 128, 64, 256, 9, true>(ABC, SFAB, outputs);
return C;
}
TORCH_LIBRARY(blockscaled_cstr_p2, m) {
m.def("gemm_cstr_p2((Tensor, Tensor, Tensor)[] ABC, (Tensor, Tensor)[] SFAB, Tensor[] outputs) -> Tensor[]");
m.impl("gemm_cstr_p2", &gemm_cstr_p2);
}
"""
CUDA_SRC_3 = r"""
using namespace cute;
constexpr int MAX_TMEM_COLS = 512;
constexpr int NUM_SMS = 148;
constexpr int NUM_BLOCKS0 = 64;
constexpr int NUM_BLOCKS1 = 192;
constexpr int GEMM_BLOCK_END0 = NUM_BLOCKS0;
constexpr int GEMM_BLOCK_END1 = GEMM_BLOCK_END0 + NUM_BLOCKS1;
constexpr int NUM_BLOCKS = GEMM_BLOCK_END1;
template<int PROBLEM_SIZE>
__constant__ CUtensorMap A_tmaps[PROBLEM_SIZE];
template<int PROBLEM_SIZE>
__constant__ CUtensorMap B_tmaps[PROBLEM_SIZE];
template<int PROBLEM_SIZE>
__constant__ CUtensorMap SFA_tmaps[PROBLEM_SIZE];
template<int PROBLEM_SIZE>
__constant__ CUtensorMap SFB_tmaps[PROBLEM_SIZE];
template<int PROBLEM_SIZE>
__constant__ half* C_ptrs[PROBLEM_SIZE];
template<int PROBLEM_SIZE>
__constant__ int Ms[PROBLEM_SIZE];
template<int PROBLEM_SIZE>
__constant__ int Ns[PROBLEM_SIZE];
template <
int K,
int BLOCK_M,
int BLOCK_N,
int BLOCK_K,
int NUM_STAGES,
class SFALayout,
class SFBLayout,
bool SWAP_AB,
int PROBLEM_SIZE,
int CTA_GROUP
>
__global__
__cluster_dims__(CTA_GROUP, 1, 1)
__launch_bounds__(BLOCK_M + 2 * WARP_SIZE) // __launch_bounds__(MAX_THREADS_PER_BLOCK, MIN_BLOCKS_PER_MP)
void multi_gemm_kernel_p3() {
// https://docs.nvidia.com/cuda/parallel-thread-execution/#tcgen05-matrix-shape
static_assert(BLOCK_M == 128); // There doesnt seem to be any code for tiling block_M
// This layout is used as TMA assumes col major, this is a colasced canonical layout.
auto asmem_layout = make_layout(make_shape(_256{}, BLOCK_M, BLOCK_K / 256), make_stride(_1{}, _256{}, 256 * BLOCK_M)); // cute::LayoutLeft{} ?
// Each member of the cluster will get a cta_rank specific offset
auto bsmem_layout = make_layout(make_shape(_256{}, BLOCK_N / CTA_GROUP, BLOCK_K / 256), make_stride(_1{}, _256{}, 256 * BLOCK_N / CTA_GROUP));
// an SF atom is ((32, 4), (16, 4)), the shmem layout is ((128 * SF_ATOMS_PER_BLOCK_K, 4)) for TMA purposes
// Since we are using CU_TENSOR_MAP_SWIZZLE_NONE, the TMA atom has shape (8, 128bits) or (8, 16) bytes, so SBO below is 8 * 16
constexpr int SF_TILE_ROWS = 128;
constexpr int SF_ATOM_NUM_ELEMENTS = 128 * 4; // cosize(SF_atom)
constexpr int SF_ATOMS_PER_BLOCK_K = BLOCK_K / 16 / 4; // each atom has 16 * 4 columns
constexpr int SF_ROWS_PER_TMEM_COL = 32; // each tmem col has 32 rows - since we are using tcgen05.cp.cta_group::1.32x128b.warpx4
using Problem = std::tuple<const CUtensorMap*, const CUtensorMap*, const CUtensorMap*, const CUtensorMap*, half*, int, int, SFALayout, SFBLayout>;
auto get_problem = [&](int problem_id) -> Problem {
auto make_sf_layout = [](int rest) {
constexpr int rest_K = K / 16 / 4;
return make_layout(make_shape(make_shape(make_shape(_32{}, _4{}), rest), make_shape(make_shape(_16{}, _4{}), rest_K)),
make_stride(make_stride(make_stride(_16{}, _4{}), rest_K * 512), make_stride(make_stride(_0{}, _1{}), _512{})));
};
const CUtensorMap* A_tmap = &A_tmaps<PROBLEM_SIZE>[problem_id];
const CUtensorMap* B_tmap = &B_tmaps<PROBLEM_SIZE>[problem_id];
const CUtensorMap* SFA_tmap = &SFA_tmaps<PROBLEM_SIZE>[problem_id];
const CUtensorMap* SFB_tmap = &SFB_tmaps<PROBLEM_SIZE>[problem_id];
half* C = C_ptrs<PROBLEM_SIZE>[problem_id];
int M = Ms<PROBLEM_SIZE>[problem_id];
int N = Ns<PROBLEM_SIZE>[problem_id];
const int rest_M = DIVUP(M, SF_TILE_ROWS);
const int rest_N = DIVUP(N, SF_TILE_ROWS);
SFALayout sfa_layout = make_sf_layout(rest_M);
SFBLayout sfb_layout = make_sf_layout(rest_N);
return {A_tmap, B_tmap, SFA_tmap, SFB_tmap, C, M, N, sfa_layout, sfb_layout};
};
// CTA rank in a cluster
int cta_rank;
asm volatile("mov.b32 %0, %%cluster_ctarank;" : "=r"(cta_rank));
const int tid = threadIdx.x;
const int lane_id = tid % WARP_SIZE;
const int warp_id = tid / WARP_SIZE;
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; // BLOCK_M * BLOCK_K elements, each element is 4 bits so divide by 2 to get bytes.
constexpr int B_size = BLOCK_N * BLOCK_K / 2 / CTA_GROUP; // each CTA only loads half of B
constexpr int SFA_size = SF_ATOM_NUM_ELEMENTS * SF_ATOMS_PER_BLOCK_K; // always copy one atom - this is ATOM_SIZE_IN_BYTES * ATOMS_PER_BLOCK_K
constexpr int SFB_size = SF_ATOM_NUM_ELEMENTS * SF_ATOMS_PER_BLOCK_K;
constexpr int STAGE_SIZE = A_size + B_size + SFA_size + SFB_size;
// set up mbarriers and tmem
// we have NUM_STAGES mbars for TMA
// NUM_STAGES mbars for MMA
// 1 mbar for mainloop
#pragma nv_diag_suppress static_var_with_dynamic_init
__shared__ int64_t mbars[NUM_STAGES * 2 + 1];
__shared__ int tmem_addr[1]; // tmem address is 32-bit
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;
const int epilogue_mbar_addr = mainloop_mbar_addr + 2 * 8;
// https://docs.nvidia.com/cuda/parallel-thread-execution/#tcgen05-mma-scale-factor-a-layout-4x
// each MMA consumes:
// - (128, 64) of A -> (128, 4) of SFA -> reshaped as (32, 4', 4) -> 4 tmem columns
constexpr int SFA_tmem = BLOCK_N * 2; // Double buffer the mma output.
constexpr int SFB_tmem = SFA_tmem + 4 * (BLOCK_K / MMA_K); // we need BLOCK_K / MMA_K many scale factors for the "tiledMMA"
if (warp_id == 0 && elect_sync()) {
for (int prefetch_idx = 0; prefetch_idx < 8; ++prefetch_idx) {
asm volatile("prefetch.tensormap [%0];" :: "l"(&A_tmaps<PROBLEM_SIZE>[prefetch_idx]) : "memory");
asm volatile("prefetch.tensormap [%0];" :: "l"(&B_tmaps<PROBLEM_SIZE>[prefetch_idx]) : "memory");
asm volatile("prefetch.tensormap [%0];" :: "l"(&SFA_tmaps<PROBLEM_SIZE>[prefetch_idx]) : "memory");
asm volatile("prefetch.tensormap [%0];" :: "l"(&SFB_tmaps<PROBLEM_SIZE>[prefetch_idx]) : "memory");
}
} else if (warp_id == 1 && elect_sync()) {
for (int i = 0; i < NUM_STAGES; i++) {
mbarrier_init(tma_mbar_addr + i * 8, CTA_GROUP); // one thread in each cluster CTA reports.
mbarrier_init(mma_mbar_addr + i * 8, 1);
}
mbarrier_init(mainloop_mbar_addr, 1); // barrier for first tmem mma buffer
mbarrier_init(mainloop_mbar_addr + 8, 1); // barrier for second tmem mma buffer
mbarrier_init(epilogue_mbar_addr, 4 * CTA_GROUP);
mbarrier_init(epilogue_mbar_addr + 8, 4 * CTA_GROUP);
asm volatile("fence.mbarrier_init.release.cluster;"); // visible to async proxy
}
if constexpr (CTA_GROUP > 1) {
// visible to all threads in a cluster
asm volatile("barrier.cluster.arrive.relaxed.aligned;");
asm volatile("barrier.cluster.wait.acquire.aligned;");
}
else {
// visible to all threads in a threadblock
__syncthreads();
}
constexpr int num_iters = K / BLOCK_K;
const int bid = blockIdx.x;
auto scheduler = [&](int global_cluster_id) -> std::tuple<int, int, int, int> {
int problem_id = 0;
int local_cluster_id = global_cluster_id;
if (global_cluster_id < GEMM_BLOCK_END0 / CTA_GROUP) {
problem_id = 0;
local_cluster_id = global_cluster_id;
} else if (global_cluster_id < GEMM_BLOCK_END1 / CTA_GROUP) {
problem_id = 1;
local_cluster_id = global_cluster_id - GEMM_BLOCK_END0 / CTA_GROUP;
}
int N = Ns<PROBLEM_SIZE>[problem_id];
const int grid_n = DIVUP(N, BLOCK_N) ;
// The cluster Ids are arranged in a (grid_m / 2 , grid_n) matrix in a row major layout.
const int bid_m = local_cluster_id / grid_n * CTA_GROUP + cta_rank;
const int bid_n = local_cluster_id % grid_n;
const int off_m = bid_m * BLOCK_M;
const int cluster_off_n = bid_n * BLOCK_N; // N offset for this cluster
const int off_n = cluster_off_n + cta_rank * (BLOCK_N / CTA_GROUP); // N offset for this CTA
return {off_m, off_n, cluster_off_n, problem_id};
};
// warp-specialization
if (warp_id == NUM_WARPS - 2 && elect_sync()) {
// TMA warp
auto issue_tma = [&](int iter_k, int stage_id, int off_m, int off_n, int problem_id) {
const int mbar_addr = (tma_mbar_addr + stage_id * 8) & 0xFEFFFFFF; // CTA0's barrier
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 problem = get_problem(problem_id);
const CUtensorMap* A_tmap = std::get<0>(problem);
const CUtensorMap* B_tmap = std::get<1>(problem);
const int M = std::get<5>(problem);
const int N = std::get<6>(problem);
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;
}
// issue TMA
const int off_k = iter_k * BLOCK_K;
tma_3d_gmem2smem<CTA_GROUP>(A_smem, A_tmap, 0, off_m, off_k / 256, mbar_addr, cache_A);
tma_3d_gmem2smem<CTA_GROUP>(B_smem, B_tmap, 0, off_n, off_k / 256, mbar_addr, cache_B);
const int rest_m = off_m / 32 / 4; // M atom is (32, 4)
const int rest_n = off_n / 32 / 4; // N atom is (32, 4)
const int rest_k = off_k / 16 / 4; // The SF atom is ((32, 4), (16, 4))
const CUtensorMap* SFA_tmap = std::get<2>(problem);
const CUtensorMap* SFB_tmap = std::get<3>(problem);
SFALayout sfa_layout = std::get<7>(problem);
SFALayout sfb_layout = std::get<8>(problem);
// Divide by 8 since underlying type is INT64
int sfa_offset = crd2idx(make_coord(make_coord(_0{}, rest_m), make_coord(_0{}, rest_k)), sfa_layout) / 8;
tma_1d_gmem2smem<CTA_GROUP>(SFA_smem, SFA_tmap, sfa_offset, mbar_addr, cache_A);
int sfb_offset = crd2idx(make_coord(make_coord(_0{}, rest_n), make_coord(_0{}, rest_k)), sfb_layout) / 8;
tma_1d_gmem2smem<CTA_GROUP>(SFB_smem, SFB_tmap, sfb_offset, mbar_addr, cache_B);
// signal TMA done .shared::cluster
asm volatile("mbarrier.arrive.expect_tx.release.cta.shared::cluster.b64 _, [%0], %1;"
:: "r"(mbar_addr), "r"(STAGE_SIZE) : "memory");
};
int mma_phase = 1; // Start with phase 1 as there are no pending MMA operations
int tma_pipeline_stage = 0;
for (int this_bid = bid; this_bid < NUM_BLOCKS; this_bid += NUM_SMS) {
auto [off_m, off_n, dummy, problem_id] = scheduler(this_bid / CTA_GROUP); // divide by CTA_GROUP to get cluster id
for (int iter_k = 0; iter_k < num_iters; iter_k++) {
// wait MMA
// NB when debugging, it will crash w/o this barrier. You cannot keep issuing TMAs w/o the
// signal the mbarrier_wait(tma_mbar_addr + stage_id * 8, tma_phase); has completed.
mbarrier_wait(mma_mbar_addr + tma_pipeline_stage * 8, mma_phase);
issue_tma(iter_k, tma_pipeline_stage, off_m, off_n, problem_id);
tma_pipeline_stage = (tma_pipeline_stage + 1) % NUM_STAGES;
if (tma_pipeline_stage == 0) {
mma_phase ^= 1;
}
}
}
}
else if (warp_id == NUM_WARPS - 1) {
// allocate tmem
const int addr = static_cast<int>(__cvta_generic_to_shared(tmem_addr));
asm volatile("tcgen05.alloc.cta_group::%2.sync.aligned.shared::cta.b32 [%0], %1;" :: "r"(addr), "r"(MAX_TMEM_COLS), "n"(CTA_GROUP));
if (cta_rank == 0 && elect_sync()) {
// MMA warp
// https://docs.nvidia.com/cuda/parallel-thread-execution/#tcgen05-instruction-descriptor
// fp4 MMA doesn't support MMA_M=64. Hence, we will use MMA_M=128 and ignore the rest.
constexpr int MMA_N = BLOCK_N;
constexpr int MMA_M = 128 * CTA_GROUP;
constexpr int16_t cta_mask = (1 << CTA_GROUP) - 1;
constexpr uint32_t i_desc = (1U << 7U) // atype=E2M1
| (1U << 10U) // btype=E2M1
| ((uint32_t)MMA_N >> 3U << 17U)
| ((uint32_t)MMA_M >> 7U << 27U)
;
int tma_phase = 0;
int mma_pipeline_stage = 0;
int epilogue_phase = 1;
int mainloop_stage = 0; // 0/1 indicates where in tmem we will put the result
for (int this_bid = bid; this_bid < NUM_BLOCKS; this_bid += NUM_SMS) {
mbarrier_wait(epilogue_mbar_addr + 8 * mainloop_stage, epilogue_phase);
auto [off_m, off_n, dummy, dummy1] = scheduler(this_bid / CTA_GROUP); // divide by CTA_GROUP to get cluster id
const int M_sf_tmem_offset = (off_m % SF_TILE_ROWS) / SF_ROWS_PER_TMEM_COL;
const int N_sf_tmem_offset = (off_n % SF_TILE_ROWS) / SF_ROWS_PER_TMEM_COL;
for (int iter_k = 0; iter_k < num_iters; iter_k++) {
// wait TMA
mbarrier_wait(tma_mbar_addr + mma_pipeline_stage * 8, tma_phase);
const int A_smem = smem + mma_pipeline_stage * 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;
// 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.
auto make_desc_AB = [](int addr) -> uint64_t {
// GMMA::Layout_K_SW128_Atom<cutlass::float_e2m1_t>{} is (_8,_256):(_256,_1) so to get to the start of the next row is 256 / 2 = 128 bytes
// so to skip 8 rows, we need 8 * 128 bytes.
const int SBO = 8 * 128;
return desc_encode(addr) | (desc_encode(SBO) << 32ULL) | (1ULL << 46ULL) | (2ULL << 61ULL);
};
// no swizzling
auto make_desc_SF = [](int addr) -> uint64_t {
// Since we are using no swizzling, the atom has size (8, 16 = 128 / 8). These atoms
// are tiled (SF_ATOMS_PER_BLOCK_K = 4, 1). So the SBO is 8 * 16
const int SBO = 8 * 16;
return desc_encode(addr) | (desc_encode(SBO) << 32ULL) | (1ULL << 46ULL);
};
// tcgen05.cp -> tcgen05.mma should be pipelined correctly per PTX doc
// https://docs.nvidia.com/cuda/parallel-thread-execution/#tcgen05-memory-consistency-model-pipelined-instructions
// cutlass issues all of smem->tmem BEFORE mma
// https://github.com/NVIDIA/cutlass/blob/v4.3.2/include/cutlass/gemm/collective/sm100_blockscaled_mma_warpspecialized.hpp#L1013-L1016
constexpr uint64_t SF_desc = make_desc_SF(0);
// The below should really be SF_desc + matrix-descriptor-encode(x) , where matrix-descriptor-encode(x) = (x & 0x3FFFF) >> 4
// but (x & 0x3FFFF) >> 4 simplifies to x >> 4.
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++) {
// The below should really be SFA_desc + (uint64_t)k * matrix-descriptor-encode(512ULL); where matrix-descriptor-encode(x) = (x & 0x3FFFF) >> 4
// but (x & 0x3FFFF) >> 4 simplifies to x >> 4. So this is SFA_desc + (uint64_t)k * (512ULL >> 4ULL)
// So we are copying 512 bytes chunks , so adjust the shmem desc accordingly.
uint64_t sfa_desc = SFA_desc + (uint64_t)k * (512ULL >> 4ULL); // 4 columns, 512 bytes of 128x4 / 32x4x4
uint64_t sfb_desc = SFB_desc + (uint64_t)k * (512ULL >> 4ULL);
// tmem addresses
// Tensor Memory addresses are 32-bit wide and specify two components.
// Lane index
// Column index
// The layout is as follows:
// 31 16
// 15 0
// Lane index
// Column index
// We are using tcgen05.cp.cta_group::1.32x128b.warpx4
// See https://docs.nvidia.com/cuda/parallel-thread-execution/#tcgen05-mma-scale-factor-a-layout-4x
// so the slices ((32, 4), (16, 4)) slices of the scale factors are in the 32 rows (lanes) of tmem
tcgen05_cp_nvfp4<CTA_GROUP>(SFA_tmem + k * 4, sfa_desc);
tcgen05_cp_nvfp4<CTA_GROUP>(SFB_tmem + k * 4, sfb_desc);
}
// Compare to tiled_mma
// k1 selects the (BLOCK_M, 256) tile.
// k2 selects the (BLOCK_M, 64) tile, whose rows are swizzled.
// NOTE: this doesn't work with BLOCK_N=32, since apparently tcgen05.mma requires SFB_tmem
// to have 2-column (8-byte) alignment (looks like not documented).
// HOWEVER: works with a cluster with BLOCK_N = 64 - maybe only the issuing thread has the alignment
// requirement.
for (int k1 = 0; k1 < BLOCK_K / 256; k1++) { // BLOCK_K == 256 so this loop only executes once
for (int k2 = 0; k2 < 256 / MMA_K; k2++) { // this is the inner loop needed since MMA_K is in general < 256 = BLOCK_K
// uint64_t a_desc = make_desc_AB(A_smem + k1 * BLOCK_M * 128 + k2 * 32);
// crd2idx(make_coord(k2 * MMA_K, 0, k1), asmem_layout) is in elements, divide by 2 to get bytes
uint64_t a_desc = make_desc_AB(A_smem + crd2idx(make_coord(k2 * MMA_K, _0{}, k1), asmem_layout) / 2);
uint64_t b_desc = make_desc_AB(B_smem + crd2idx(make_coord(k2 * MMA_K, _0{}, k1), bsmem_layout) / 2);
int k_sf = k1 * 4 + k2; // the scale factors are in blocks of 4. This is the block idx
// k_sf is mulitplied by 4 to get us to the actual column.
const int scale_A_tmem = SFA_tmem + k_sf * 4 + M_sf_tmem_offset;
const int scale_B_tmem = SFB_tmem + k_sf * 4 + N_sf_tmem_offset;
const int enable_input_d = (k1 == 0 && k2 == 0) ? iter_k : 1;
tcgen05_mma_nvfp4<CTA_GROUP>(/*dst tmem*/ BLOCK_N * mainloop_stage, a_desc, b_desc, i_desc, scale_A_tmem, scale_B_tmem, enable_input_d);
} // for k2
} // for k1
// signal MMA done
// @@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@
// tcgen05_commit_mcast does not work if CTA_GROUP = 1
// tcgen05_commit_mcast<CTA_GROUP>(mma_mbar_addr + mma_pipeline_stage * 8, cta_mask); // signal MMA done
// @@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@
asm volatile("tcgen05.commit.cta_group::%2.mbarrier::arrive::one.shared::cluster.multicast::cluster.b64 [%0], %1;"
:: "r"(mma_mbar_addr + mma_pipeline_stage * 8), "h"(cta_mask), "n"(CTA_GROUP) : "memory");
mma_pipeline_stage = (mma_pipeline_stage + 1) % NUM_STAGES;
if (mma_pipeline_stage == 0) {
tma_phase ^= 1;
}
}
// signal mainloop done
// tcgen05_commit_mcast<CTA_GROUP>(mainloop_mbar_addr + mainloop_stage * 8, cta_mask); // signal mainloop done
asm volatile("tcgen05.commit.cta_group::%2.mbarrier::arrive::one.shared::cluster.multicast::cluster.b64 [%0], %1;"
:: "r"(mainloop_mbar_addr + mainloop_stage * 8), "h"(cta_mask), "n"(CTA_GROUP) : "memory");
mainloop_stage = (mainloop_stage + 1) % 2;
if (mainloop_stage == 0) {
epilogue_phase ^= 1;
}
} // for bid
} // cta_rank == 0 && elect_sync()
} // warp_id == NUM_WARPS - 1
else if (tid < BLOCK_M) {
// epilogue warps
int mainloop_phase = 0;
int mainloop_stage = 0; // 0/1 indicates where in tmem we will put the result
for (int this_bid = bid; this_bid < NUM_BLOCKS; this_bid += NUM_SMS) {
auto result = scheduler(this_bid / CTA_GROUP); // divide by CTA_GROUP to get cluster id
int off_m = std::get<0>(result);
int cluster_off_n = std::get<2>(result);
int problem_id = std::get<3>(result);
auto problem = get_problem(problem_id);
const int M = std::get<5>(problem);
const int N = std::get<6>(problem);
half* C_ptr = std::get<4>(problem);
// wait mainloop
mbarrier_wait(mainloop_mbar_addr + mainloop_stage * 8, mainloop_phase);
asm volatile("tcgen05.fence::after_thread_sync;");
int tmem_buffer_offset = BLOCK_N * mainloop_stage;
auto epilogue_M_major = [&]() {
// C is M-major
constexpr int WIDTH = std::min(BLOCK_N, 64); // using 128 might be slower
for (int n = 0; n < (BLOCK_N + WIDTH - 1)/ WIDTH; n++) {
float tmp[WIDTH]; // if WIDTH=128, we are using 128 registers here
if constexpr (WIDTH == 128) tcgen05_ld_32x32bx128(tmp, warp_id * 32, n * WIDTH + tmem_buffer_offset);
if constexpr (WIDTH == 64) tcgen05_ld_32x32bx64(tmp, warp_id * 32, n * WIDTH + tmem_buffer_offset);
if constexpr (WIDTH == 32) tcgen05_ld_32x32bx32(tmp, warp_id * 32, n * WIDTH + tmem_buffer_offset);
asm volatile("tcgen05.wait::ld.sync.aligned;");
for (int i = 0; i < WIDTH; i++) {
// NB row and col are transposed.
const int row = cluster_off_n + n * WIDTH + i;
const int col = off_m + tid;
if (row < N && col < M) {
C_ptr[row * M + col] = __float2half(tmp[i]);
}
}
}
};
auto epilogue_N_major = [&]() {
// C is N-major
for (int m = 0; m < 32 / 16; m++) {
float tmp[BLOCK_N / 2];
if constexpr (BLOCK_N == 128) tcgen05_ld_16x256bx16(tmp, warp_id * 32 + m * 16, tmem_buffer_offset);
if constexpr (BLOCK_N == 64) tcgen05_ld_16x256bx8(tmp, warp_id * 32 + m * 16, tmem_buffer_offset);
if constexpr (BLOCK_N == 32) tcgen05_ld_16x256bx4(tmp, warp_id * 32 + m * 16, tmem_buffer_offset);
asm volatile("tcgen05.wait::ld.sync.aligned;");
for (int i = 0; i < (BLOCK_N + 7)/ 8; i++) {
// TODO replace this with cute tensors.
const int row = off_m + warp_id * 32 + m * 16 + lane_id / 4;
const int col = cluster_off_n + i * 8 + (lane_id % 4) * 2;
if (row < M) {
reinterpret_cast<half2 *>(C_ptr + (row + 0) * N + col)[0] = __float22half2_rn({tmp[i * 4 + 0], tmp[i * 4 + 1]});
}
if (row + 8 < M) {
reinterpret_cast<half2 *>(C_ptr + (row + 8) * N + col)[0] = __float22half2_rn({tmp[i * 4 + 2], tmp[i * 4 + 3]});
}
}
}
};
if constexpr (SWAP_AB) {
epilogue_M_major();
} else {
epilogue_N_major();
}
if (elect_sync()) {
// // signal Epilogue done
// constexpr int16_t cta_mask = 3;
// tcgen05_commit_mcast<CTA_GROUP>(epilogue_mbar_addr + 8 * mainloop_stage, cta_mask); // signal Epilogue done
const int mbar_addr = (epilogue_mbar_addr + mainloop_stage * 8) & 0xFEFFFFFF;
asm volatile("mbarrier.arrive.release.cta.shared::cluster.b64 _, [%0];" :: "r"(mbar_addr) : "memory");
}
mainloop_stage = (mainloop_stage + 1) % 2;
if (mainloop_stage == 0) {
mainloop_phase ^= 1;
}
} // for bid
} //epilogue warps
// All warps.
if constexpr (CTA_GROUP > 1) {
// Dont exit w/o the peer CTA.
asm volatile("barrier.cluster.arrive.release.aligned;");
asm volatile("barrier.cluster.wait.acquire.aligned;");
} else {
__syncthreads(); // all threads finish reading data from tmem
}
// deallocate tmem. tmem address should be 0.
if (warp_id == 0) {
asm volatile("tcgen05.dealloc.cta_group::%2.sync.aligned.b32 %0, %1;" :: "r"(0), "r"(MAX_TMEM_COLS), "n"(CTA_GROUP));
}
}
template <
int K,
int BLOCK_M,
int BLOCK_N,
int BLOCK_K,
int NUM_STAGES,
bool SWAP_AB
>
std::vector<at::Tensor> launch_gemm_cstr_p3(
c10::List<std::tuple<at::Tensor, at::Tensor, at::Tensor>> ABC,
c10::List<std::tuple<at::Tensor, at::Tensor>> SFAB,
c10::List<at::Tensor> outputs) {
constexpr int CTA_GROUP = 2;
static_assert(BLOCK_K % 256 == 0);
constexpr int rest_K = K / 16 / 4;
constexpr int grid = NUM_SMS;
constexpr int tb_size = BLOCK_M + 2 * WARP_SIZE;
constexpr int AB_size = (BLOCK_M + BLOCK_N / CTA_GROUP) * (BLOCK_K / 2);
constexpr int SFAB_size = 128 * (BLOCK_K / 16) * 2;
constexpr int smem_size = (AB_size + SFAB_size) * NUM_STAGES;
// This layout is used as TMA assumes col major, this is a colasced canonical layout.
auto layout_smem_A = make_layout(make_shape(256, BLOCK_M, BLOCK_K / 256), make_stride(1, 256, 256 * BLOCK_M)); // cute::LayoutLeft{} ?
// Each member of the cluster will get a cta_rank specific offset
auto layout_smem_B = make_layout(make_shape(256, BLOCK_N / CTA_GROUP, BLOCK_K / 256), make_stride(1, 256, 256 * BLOCK_N / CTA_GROUP));
constexpr int SF_TILE_ROWS = 128;
constexpr int SF_ATOM_NUM_ELEMENTS = 128 * 4; // cosize(SF_atom)
constexpr int SF_ATOMS_PER_BLOCK_K = BLOCK_K / 16 / 4; // each atom has 16 * 4 columns
auto make_sf_layout = [&](int rest) {
return make_layout(make_shape(make_shape(make_shape(_32{}, _4{}), rest), make_shape(make_shape(_16{}, _4{}), rest_K)),
make_stride(make_stride(make_stride(_16{}, _4{}), rest_K * 512), make_stride(make_stride(_0{}, _1{}), _512{})));
};
// This is for the type only
auto layout_sfa = make_sf_layout(16);
auto layout_sfb = make_sf_layout(16);
auto populate_kernel_config = [&ABC, &SFAB, &outputs](CUtensorMap* A_tmaps, CUtensorMap* B_tmaps, CUtensorMap* SFA_tmaps, CUtensorMap* SFB_tmaps,
half** C_ptrs, int* Ms, int* Ns, auto problem_size_const) {
constexpr int PROBLEM_SIZE = decltype(problem_size_const)::value;
for (int workItem = 0; workItem < PROBLEM_SIZE; ++workItem) {
std::tuple<at::Tensor, at::Tensor, at::Tensor> abc_tuple = ABC[workItem];
auto A = std::get<0>(abc_tuple);
auto B = std::get<1>(abc_tuple);
std::tuple<at::Tensor, at::Tensor> sfab_tuple = SFAB[workItem];
auto SFA = std::get<0>(sfab_tuple);
auto SFB = std::get<1>(sfab_tuple);
at::Tensor C = outputs[workItem];
const int M = A.size(0);
const int N = B.size(0);
auto A_ptr = reinterpret_cast<const char *>(A.data_ptr());
auto B_ptr = reinterpret_cast<const char *>(B.data_ptr());
auto SFA_ptr = reinterpret_cast<const char *>(SFA.data_ptr());
auto SFB_ptr = reinterpret_cast<const char *>(SFB.data_ptr());
C_ptrs[workItem] = reinterpret_cast<half *>(C.data_ptr());
int new_M = M;
int new_N = N;
if constexpr (SWAP_AB) {
std::swap(A_ptr, B_ptr);
std::swap(SFA_ptr, SFB_ptr);
std::swap(new_M, new_N);
}
Ms[workItem] = new_M;
Ns[workItem] = new_N;
init_AB_tmap(&A_tmaps[workItem], A_ptr, new_M, K, BLOCK_M, BLOCK_K);
init_AB_tmap(&B_tmaps[workItem], B_ptr, new_N, K, BLOCK_N / CTA_GROUP, BLOCK_K);
// CUtensorMap SFA_tmap0, SFB_tmap0;
// SF Atom is ((32, 4), (16, 4))
const int rest_M = DIVUP(new_M, SF_TILE_ROWS);
const int rest_N = DIVUP(new_N, SF_TILE_ROWS);
auto new_M_padded = SF_TILE_ROWS * rest_M; // 128 is the number of rows in the atom.
auto new_N_padded = SF_TILE_ROWS * rest_N;
init_SF_tmap(&SFA_tmaps[workItem], SFA_ptr, new_M_padded * K / SF_BLOCK_SIZE, SF_ATOM_NUM_ELEMENTS * SF_ATOMS_PER_BLOCK_K);
init_SF_tmap(&SFB_tmaps[workItem], SFB_ptr, new_N_padded * K / SF_BLOCK_SIZE, SF_ATOM_NUM_ELEMENTS * SF_ATOMS_PER_BLOCK_K);
}
};
constexpr int PROBLEM_SIZE = 2;
CUtensorMap host_A_tmaps[PROBLEM_SIZE];
CUtensorMap host_B_tmaps[PROBLEM_SIZE];
CUtensorMap host_SFA_tmaps[PROBLEM_SIZE];
CUtensorMap host_SFB_tmaps[PROBLEM_SIZE];
half* host_C_ptrs[PROBLEM_SIZE];
int host_Ms[PROBLEM_SIZE];
int host_Ns[PROBLEM_SIZE];
populate_kernel_config(host_A_tmaps, host_B_tmaps, host_SFA_tmaps, host_SFB_tmaps, host_C_ptrs, host_Ms, host_Ns, std::integral_constant<int, PROBLEM_SIZE>{});
cudaMemcpyToSymbol(A_tmaps<PROBLEM_SIZE>, host_A_tmaps, PROBLEM_SIZE * sizeof(CUtensorMap));
cudaMemcpyToSymbol(B_tmaps<PROBLEM_SIZE>, host_B_tmaps, PROBLEM_SIZE * sizeof(CUtensorMap));
cudaMemcpyToSymbol(SFA_tmaps<PROBLEM_SIZE>, host_SFA_tmaps, PROBLEM_SIZE * sizeof(CUtensorMap));
cudaMemcpyToSymbol(SFB_tmaps<PROBLEM_SIZE>, host_SFB_tmaps, PROBLEM_SIZE * sizeof(CUtensorMap));
cudaMemcpyToSymbol(C_ptrs<PROBLEM_SIZE>, host_C_ptrs, PROBLEM_SIZE * sizeof(half*));
cudaMemcpyToSymbol(Ms<PROBLEM_SIZE>, host_Ms, PROBLEM_SIZE * sizeof(int));
cudaMemcpyToSymbol(Ns<PROBLEM_SIZE>, host_Ns, PROBLEM_SIZE * sizeof(int));
auto this_kernel = multi_gemm_kernel_p3<K, BLOCK_M, BLOCK_N, BLOCK_K, NUM_STAGES,
decltype(layout_sfa), decltype(layout_sfb),
SWAP_AB, PROBLEM_SIZE, CTA_GROUP>;
if (smem_size > 48'000)
cudaFuncSetAttribute(this_kernel, cudaFuncAttributeMaxDynamicSharedMemorySize, smem_size);
this_kernel<<<grid, tb_size, smem_size>>>();
return {outputs[0], outputs[1]};
}
template std::vector<at::Tensor> launch_gemm_cstr_p3<1536, 128, 64, 256, 9, true>(
c10::List<std::tuple<at::Tensor, at::Tensor, at::Tensor>> ABC,
c10::List<std::tuple<at::Tensor, at::Tensor>> SFAB,
c10::List<at::Tensor> outputs);
"""
CPP_SRC_3 = r"""
#include <vector>
#include <torch/script.h>
template <
int K,
int BLOCK_M,
int BLOCK_N,
int BLOCK_K,
int NUM_STAGES,
bool SWAP_AB
>
std::vector<at::Tensor> launch_gemm_cstr_p3(
c10::List<std::tuple<at::Tensor, at::Tensor, at::Tensor>> ABC,
c10::List<std::tuple<at::Tensor, at::Tensor>> SFAB,
c10::List<at::Tensor> outputs);
std::vector<at::Tensor> gemm_cstr_p3(
c10::List<std::tuple<at::Tensor, at::Tensor, at::Tensor>> ABC,
c10::List<std::tuple<at::Tensor, at::Tensor>> SFAB,
c10::List<at::Tensor> outputs) {
auto C = launch_gemm_cstr_p3<1536, 128, 64, 256, 9, true>(ABC, SFAB, outputs);
return C;
}
TORCH_LIBRARY(blockscaled_cstr_p3, m) {
m.def("gemm_cstr_p3((Tensor, Tensor, Tensor)[] ABC, (Tensor, Tensor)[] SFAB, Tensor[] outputs) -> Tensor[]");
m.impl("gemm_cstr_p3", &gemm_cstr_p3);
}
"""
ON_VERDA = False
optional_args = {"extra_include_paths": ['/root/cutlass/include/', '/root/cutlass/tools/util/include/',]} if ON_VERDA else {}
module = load_inline(
"some_module_name_0",
cpp_sources=CPP_SRC_0,
cuda_sources=COMMON_CU + CUDA_SRC_0,
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",
],
extra_ldflags=["-lcuda"],
**optional_args
)
module = load_inline(
"some_module_name_1",
cpp_sources=CPP_SRC_1,
cuda_sources=COMMON_CU + CUDA_SRC_1,
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",
],
extra_ldflags=["-lcuda"],
**optional_args
)
module = load_inline(
"some_module_name_2",
cpp_sources=CPP_SRC_2,
cuda_sources=COMMON_CU + CUDA_SRC_2,
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",
],
extra_ldflags=["-lcuda"],
**optional_args
)
module = load_inline(
"some_module_name_3",
cpp_sources=CPP_SRC_3,
cuda_sources=COMMON_CU + CUDA_SRC_3,
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",
],
extra_ldflags=["-lcuda"],
**optional_args
)
################################################################################
########################## Test custom kernel ##################################
################################################################################
# Scaling factor vector size
sf_vec_size = 16
# Helper function for ceiling division
def ceil_div(a, b):
return (a + b - 1) // b
# Helper function to convert scale factor tensor to blocked format
# Helper function to convert scale factor tensor to blocked format
def to_blocked(input_matrix):
rows, cols = input_matrix.shape
# Please ensure rows and cols are multiples of 128 and 4 respectively
n_row_blocks = ceil_div(rows, 128)
n_col_blocks = ceil_div(cols, 4)
padded_rows = n_row_blocks * 128
padded_cols = n_col_blocks * 4
# Pad the input matrix if necessary
if padded_rows != rows or padded_cols != cols:
padded = torch.nn.functional.pad(
input_matrix,
(0, padded_cols - cols, 0, padded_rows - rows),
mode="constant",
value=0,
)
else:
padded = input_matrix
blocks = padded.view(n_row_blocks, 128, n_col_blocks, 4).permute(0, 2, 1, 3)
rearranged = blocks.reshape(-1, 4, 32, 4).transpose(1, 2).reshape(-1, 32, 16)
return rearranged.flatten()
def ref_kernel(
data: input_t,
) -> output_t:
"""
PyTorch reference implementation of NVFP4 block-scaled group GEMM.
"""
abc_tensors, sfasfb_tensors, _, problem_sizes = data
result_tensors = []
for i, (
(a_ref, b_ref, c_ref),
(sfa_ref, sfb_ref),
(m, n, k, l),
) in enumerate(
zip(
abc_tensors,
sfasfb_tensors,
problem_sizes,
)
):
for l_idx in range(l):
# Convert the scale factor tensor to blocked format
scale_a = to_blocked(sfa_ref[:, :, l_idx])
scale_b = to_blocked(sfb_ref[:, :, l_idx])
# (m, k) @ (n, k).T -> (m, n)
res = torch._scaled_mm(
a_ref[:, :, l_idx].view(torch.float4_e2m1fn_x2),
b_ref[:, :, l_idx].transpose(0, 1).view(torch.float4_e2m1fn_x2),
scale_a.cuda(),
scale_b.cuda(),
bias=None,
out_dtype=torch.float16,
)
c_ref[:, :, l_idx] = res
result_tensors.append((c_ref))
return result_tensors
# Helper function to prepare the scale factor tensors for both reference
# kernel and customize kernel. The customized data layout can be found in:
# https://docs.nvidia.com/cuda/cublas/index.html?highlight=fp4#d-block-scaling-factors-layout
def create_reordered_scale_factor_tensor(l, mn, k, ref_f8_tensor):
sf_k = ceil_div(k, sf_vec_size)
atom_m = (32, 4)
atom_k = 4
mma_shape = (
l, # batch size
ceil_div(mn, atom_m[0] * atom_m[1]),
ceil_div(sf_k, atom_k),
atom_m[0],
atom_m[1],
atom_k,
)
# Create the reordered scale factor tensor (32, 4, rest_m, 4, rest_k, l) on GPU.
mma_permute_order = (3, 4, 1, 5, 2, 0)
# Generate a random int8 tensor, then convert to float8_e4m3fn
rand_int_tensor = torch.randint(1, 3, mma_shape, dtype=torch.int8, device='cuda')
reordered_f8_tensor = rand_int_tensor.to(dtype=torch.float8_e4m3fn)
# Permute according to mma_permute_order
reordered_f8_tensor = reordered_f8_tensor.permute(*mma_permute_order)
# Move ref_f8_tensor to GPU if not already there
if ref_f8_tensor.device.type == 'cpu':
ref_f8_tensor = ref_f8_tensor.cuda()
# GPU-side vectorized reordering (replaces slow CPU nested loops)
# Create index grids for all dimensions
i_idx = torch.arange(mn, device='cuda')
j_idx = torch.arange(sf_k, device='cuda')
b_idx = torch.arange(l, device='cuda')
# Create meshgrid for all combinations of (i, j, b)
i_grid, j_grid, b_grid = torch.meshgrid(i_idx, j_idx, b_idx, indexing='ij')
# Calculate target indices in vectorized manner
mm = i_grid // (atom_m[0] * atom_m[1])
mm32 = i_grid % atom_m[0]
mm4 = (i_grid % 128) // atom_m[0]
kk = j_grid // atom_k
kk4 = j_grid % atom_k
# Perform the reordering with advanced indexing (all on GPU)
reordered_f8_tensor[mm32, mm4, mm, kk4, kk, b_grid] = ref_f8_tensor[i_grid, j_grid, b_grid]
return reordered_f8_tensor
def _create_fp4_tensors(l, mn, k):
# generate uint8 tensor, then convert to float4e2m1fn_x2 data type
# generate all bit patterns
ref_i8 = torch.randint(255, size=(l, mn, k // 2), dtype=torch.uint8, device="cuda") # * 0 + 34 # Remove comment to make inputs one 34 = b'0010'0010
# for each nibble, only keep the sign bit and 2 LSBs
# the possible values are [-1.5, -1, -0.5, 0, +0.5, +1, +1.5]
ref_i8 = ref_i8 & 0b1011_1011
return ref_i8.permute(1, 2, 0).view(torch.float4_e2m1fn_x2)
def generate_input(
m: tuple,
n: tuple,
k: tuple,
g: int,
seed: int,
):
"""
Generate input tensors for NVFP4 block-scaled group GEMM.
Each group can have different m, n, k, l.
Args:
problem_sizes: List of tuples (m, n, k, l) for each problem
m: Number of rows in matrix A
n: Number of columns in matrix B
k: Number of columns in A and rows of B
l: Batch size, always is 1
groups: Number of groups
seed: Random seed for reproducibility
Returns:
Tuple of (list(tuple(a, b, c)), list(tuple(sfa, sfb)), list(tuple(sfa_reordered, sfb_reordered)), list(tuple(m, n, k, l))) where each group has its own a, b, c, sfa, sfb.
a: [m, k, l] - Input matrix in torch.float4e2m1fn_x2 data type
b: [n, k, l] - Input matrix in torch.float4e2m1fn_x2 data type
sfa: [m, k // 16, l] - Input scale factors in torch.float8e4m3fn data type
sfb: [n, k // 16, l] - Input scale factors in torch.float8e4m3fn data type
sfa_reordered: [32, 4, rest_m, 4, rest_k, l] - Input scale factors in torch.float8e4m3fn data type
sfb_reordered: [32, 4, rest_n, 4, rest_k, l] - Input scale factors in torch.float8e4m3fn data type
c: [m, n, l] - Output matrix in torch.float16 data type
"""
torch.manual_seed(seed)
abc_tensors = []
sfasfb_tensors = []
sfasfb_reordered_tensors = []
problem_sizes = []
l = 1
# Generate a, b, c, sfa, sfb tensors for all groups
for group_idx in range(g):
mi = m[group_idx]
ni = n[group_idx]
ki = k[group_idx]
a_ref = _create_fp4_tensors(l, mi, ki)
b_ref = _create_fp4_tensors(l, ni, ki)
c_ref = torch.randn((l, mi, ni), dtype=torch.float16, device="cuda").permute(
1, 2, 0
)
sf_k = ceil_div(ki, sf_vec_size)
sfa_ref_cpu_random = torch.randint(
1, 3, (l, mi, sf_k), dtype=torch.int8
) # * 0 + 1 # Remove comment to make sfs one
sfa_ref_cpu = sfa_ref_cpu_random.to(dtype=torch.float8_e4m3fn).permute(1, 2, 0)
sfb_ref_cpu_random = torch.randint(
1, 3, (l, ni, sf_k), dtype=torch.int8
) # * 0 + 1 # Remove comment to make sfs one
sfb_ref_cpu = sfb_ref_cpu_random.to(dtype=torch.float8_e4m3fn).permute(1, 2, 0)
sfa_reordered = create_reordered_scale_factor_tensor(l, mi, ki, sfa_ref_cpu)
sfb_reordered = create_reordered_scale_factor_tensor(l, ni, ki, sfb_ref_cpu)
abc_tensors.append((a_ref, b_ref, c_ref))
sfasfb_tensors.append((sfa_ref_cpu, sfb_ref_cpu))
sfasfb_reordered_tensors.append((sfa_reordered, sfb_reordered))
problem_sizes.append((mi, ni, ki, l))
return (abc_tensors, sfasfb_tensors, sfasfb_reordered_tensors, problem_sizes)
################################################################################
########################## End Test custom kernel ##############################
################################################################################
start = 0
BIG_BUFFER = torch.zeros(int(2e10), dtype=torch.float16, device="cuda")
def allocate(c: torch.Tensor):
if not USE_ALLOCATE:
return c
global start
end = start + c.numel()
buf = BIG_BUFFER[start:end].as_strided(c.shape, c.stride())
start = end
return buf
USE_ALLOCATE = False
result_tensor_func = allocate if USE_ALLOCATE else lambda x: x
def custom_kernel(data: input_t) -> output_t:
"""
Reference implementation of block-scale fp4 group gemm
Args:
data: list of tuples (abc_tensors, sfasfb_tensors, sfasfb_reordered_tensors, problem_sizes) where:
abc_tensors: list of tuples (a, b, c) where
a is torch.Tensor[float4e2m1fn_x2] of shape [m, k // 2, l]
b is torch.Tensor[float4e2m1fn_x2] of shape [n, k // 2, l]
c is torch.Tensor[float16] of shape [m, n, l]
sfasfb_tensors: list of tuples (sfa, sfb) where
sfa is torch.Tensor[float8_e4m3fnuz] of shape [m, k // 16, l]
sfb is torch.Tensor[float8_e4m3fnuz] of shape [n, k // 16, l]
sfasfb_reordered_tensors: list of tuples (sfa_reordered, sfb_reordered) where
sfa_reordered is torch.Tensor[float8_e4m3fnuz] of shape [32, 4, rest_m, 4, rest_k, l]
sfb_reordered is torch.Tensor[float8_e4m3fnuz] of shape [32, 4, rest_n, 4, rest_k, l]
problem_sizes: list of tuples (m, n, k, l)
each group has its own a, b, c, sfa, sfb with different m, n, k, l problem sizes
l should always be 1 for each group.
Returns:
list of tuples (c) where c is torch.Tensor[float16] of shape [m, n, l]
"""
abc_tensors, sfasfb_tensors, sfasfb_reordered_tensors, problem_sizes = data
output_tensors = []
for item in abc_tensors:
_, _, c_ref = item
output_tensors.append(result_tensor_func(c_ref))
if abc_tensors[0][0].shape[1] == 7168 // 2:
return torch.ops.blockscaled_cstr_p0.gemm_cstr_p0(abc_tensors, sfasfb_reordered_tensors, output_tensors)
elif abc_tensors[0][0].shape[1] == 2048 // 2:
return torch.ops.blockscaled_cstr_p1.gemm_cstr_p1(abc_tensors, sfasfb_reordered_tensors, output_tensors)
elif abc_tensors[0][0].shape[1] == 4096 // 2:
return torch.ops.blockscaled_cstr_p2.gemm_cstr_p2(abc_tensors, sfasfb_reordered_tensors, output_tensors)
elif abc_tensors[0][0].shape[1] == 1536 // 2:
return torch.ops.blockscaled_cstr_p3.gemm_cstr_p3(abc_tensors, sfasfb_reordered_tensors, output_tensors)
return ref_kernel(data)
if __name__ == '__main__':
benchmarks = [[8, [80, 176, 128, 72, 64, 248, 96, 160], [4096, 4096, 4096, 4096, 4096, 4096, 4096, 4096], [7168, 7168, 7168, 7168, 7168, 7168, 7168, 7168]],\
[8, [40, 76, 168, 72, 164, 148, 196, 160], [7168, 7168, 7168, 7168, 7168, 7168, 7168, 7168], [2048, 2048, 2048, 2048, 2048, 2048, 2048, 2048]],\
[2, [192, 320], [3072, 3072], [4096, 4096]],\
[2, [128, 384], [4096, 4096], [1536, 1536]]]
for benchmark in benchmarks:
G, M, N, K = benchmark
print(f"{M=}, {N=}, {K=}")
input_data = generate_input(M, N, K, G, 1234)
raw_results = custom_kernel(input_data)
ref_results = ref_kernel(input_data)
print(raw_results[0] - ref_results[0])scrolls · 3661 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