submission 505935
Nareg · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 1304 lines, June 9 Researcher Reciprocity License v1.0.
v0.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-nvfp4-group-gemm-505935?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:95382edbc09d0f58a84100ab58df8198732a7d03710d066c6a6ee845cc6cfaaa
license declaredunknown
license concludedunknown
authorsNareg
imported2026-08-15
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
cluster
__global__ void __cluster_dims__(CLUSTER_SIZE, 1, 1) nvfp4_group_gemm_kernel(const __grid_constant__ GroupDescs group_descs, const __grid_constant__ CUtensorMap tmap_a_temp,fp4
Need TD_SMEM_M * TD_SMEM_K * sizeof(nvfp4) bytes for a_smemmbarrier
"tcgen05.commit.cta_group::1.mbarrier::arrive::one.shared::cluster.b64 [%0];"num-warps = 7
constexpr int NUM_WARPS = 7;persistent-kernel
for (int tile_idx = blockIdx.x; tile_idx < total_tiles; tile_idx += gridDim.x) {shared-memory
__device__ uint64_t inline make_smem_desc(int smem_addr) {stages = 1
__shared__ int N_vals[NK_VAR ? TILE_DESC_PIPE_STAGES : 1];tcgen05
"tcgen05.commit.cta_group::1.mbarrier::arrive::one.shared::cluster.b64 [%0];"tma
"cp.async.bulk.shared::cta.global.mbarrier::complete_tx::bytes.L2::cache_hint [%0], [%1], %2, [%3], %4;"vector-width = half2
reinterpret_cast<half2 *>(c_smem + (C_CHUNK_SMEM_TILESZ * buf) + m_offset*OUT_N_CHUNK + n_offset)[0] = __float22half2_rn({results[i * 4], results[i * 4 + 1]});Kernel source
v0.py1304 lines
#!POPCORN leaderboard nvfp4_group_gemm
import torch
from torch.utils.cpp_extension import load_inline
from typing import List
from task import input_t, output_t
from utils import make_match_reference
"""
TMA multi-cast
"""
nvfp4_group_gemm_cuda_source = """
#include <torch/library.h>
#include <ATen/core/Tensor.h>
#include <cudaTypedefs.h> // PFN_cuTensorMapEncodeTiled, CUtensorMap
#include <cuda_fp16.h>
#include <chrono>
/*
Warp specialization and pipelining
*/
/*
Notes:
- TD stands for Tile Dimension
Assumptions:
1) Problem shape is divisible by CTA and SMEM Tile shapes (no tail cases)
2) We assume TD_SMEM_M/N == TD_MMA_M/N, since TMEM can only store one MMA tile worth of results at a time
*/
enum class CacheHintSm100 : uint64_t {
EVICT_NORMAL = 0x1000000000000000,
EVICT_FIRST = 0x12F0000000000000,
EVICT_LAST = 0x14F0000000000000,
};
#define CEIL_DIV(x, y) (((x) + (y) - 1) / (y))
__device__ void inline tcgen05_commit(const int mbar_addr) {
asm volatile(
"tcgen05.commit.cta_group::1.mbarrier::arrive::one.shared::cluster.b64 [%0];"
:
: "r"(mbar_addr)
: "memory"
);
}
template<uint16_t CLUSTER_MASK>
__device__ void inline tcgen05_commit_multicast(const int mbar_addr) {
asm volatile(
"tcgen05.commit.cta_group::1.mbarrier::arrive::one.shared::cluster.multicast::cluster.b64 [%0], %1;"
:
: "r"(mbar_addr), "h" (CLUSTER_MASK)
: "memory"
);
}
__device__ void mbar_wait(const int mbar_addr, const int phase) {
uint32_t ticks = 0x989680; // expiration date for try wait to re-try, from CUTLASS
asm volatile(
"{\\n"
".reg .pred P1;\\n"
"LAB_WAIT:\\n"
"mbarrier.try_wait.parity.shared::cta.b64 P1, [%0], %1, %2;\\n" // Acquire semantics assumed here
"@P1 bra.uni DONE;\\n" // Add .uni here because there won't be warp divergence
"bra.uni LAB_WAIT;\\n"
"DONE:\\n"
"}"
:
: "r"(mbar_addr), "r"(phase), "r"(ticks)
);
}
__device__ inline void mbar_arrive_expect(const int mbar_addr, const int size) {
asm volatile(
"mbarrier.arrive.expect_tx.release.cta.shared::cta.b64 _, [%0], %1;"
:
: "r"(mbar_addr), "r"(size)
: "memory"
);
}
__device__ inline void mbar_arrive(const int mbar_addr, const int count) {
asm volatile(
"mbarrier.arrive.release.cta.shared::cta.b64 _, [%0], %1;"
:
: "r"(mbar_addr), "r"(count)
: "memory"
);
}
__device__ inline void mbar_init(const int mbar_addr, const int count) {
asm volatile(
"mbarrier.init.shared::cta.b64 [%0], %1;"
:
: "r"(mbar_addr), "r"(count)
);
}
__device__ uint32_t inline elect_one_sync() {
uint32_t pred = 0;
asm volatile(
"{\\n"
".reg .pred %%px;\\n"
" elect.sync _|%%px, %1;\\n"
"@%%px mov.s32 %0, 1;\\n"
"}"
: "+r"(pred)
: "r"(0xFFFFFFFF)
);
return pred;
}
__device__ void inline tcgen05_1dtma_g2s_sf(int dst, const void *src, int size, int mbar_addr, CacheHintSm100 cache_policy) {
asm volatile(
"cp.async.bulk.shared::cta.global.mbarrier::complete_tx::bytes.L2::cache_hint [%0], [%1], %2, [%3], %4;"
:
: "r"(dst), "l"(src), "r"(size), "r"(mbar_addr), "l"(cache_policy));
}
template<int CTA_GROUP>
__device__ void inline tcgen05_3dtma_g2s_ab(int dst_smem, const void *tmap_ptr, int mn_off, int k_off_coremat, int mbar_addr, CacheHintSm100 cache_policy) {
asm volatile (
"cp.async.bulk.tensor.3d.shared::cluster.global.mbarrier::complete_tx::bytes.cta_group::%0.L2::cache_hint [%1], [%2, {%3, %4, %5}], [%6], %7;"
:
: "n"(CTA_GROUP), "r"(dst_smem), "l"(tmap_ptr), "r"(0), "r"(mn_off), "r"(k_off_coremat), "r"(mbar_addr), "l"(cache_policy)
);
}
template<int CTA_GROUP, uint16_t CLUSTER_MASK>
__device__ void inline tcgen05_3dtma_g2s_ab_multicast(int dst_smem, const void *tmap_ptr, int mn_off, int k_off_coremat, int mbar_addr, CacheHintSm100 cache_policy) {
asm volatile (
"cp.async.bulk.tensor.3d.shared::cluster.global.mbarrier::complete_tx::bytes.multicast::cluster.cta_group::%0.L2::cache_hint [%1], [%2, {%3, %4, %5}], [%6], %7, %8;"
:
: "n"(CTA_GROUP), "r"(dst_smem), "l"(tmap_ptr), "r"(0), "r"(mn_off), "r"(k_off_coremat), "r"(mbar_addr), "h"(CLUSTER_MASK), "l"(cache_policy)
);
}
__device__ void inline tcgen05_2dtma_s2g_c(int src_smem, const void *tmap_ptr, int m_off, int n_off, CacheHintSm100 cache_policy) {
asm volatile (
"cp.async.bulk.tensor.2d.global.shared::cta.bulk_group.L2::cache_hint [%0, {%1, %2}], [%3], %4;"
:
: "l"(tmap_ptr), "r"(n_off), "r"(m_off), "r"(src_smem), "l"(cache_policy)
);
}
PFN_cuTensorMapEncodeTiled_v12000 get_cuTensorMapEncodeTiled() {
cudaDriverEntryPointQueryResult driver_status;
void* cuTensorMapEncodeTiled_ptr = nullptr;
cudaGetDriverEntryPointByVersion("cuTensorMapEncodeTiled", &cuTensorMapEncodeTiled_ptr, 12000, cudaEnableDefault, &driver_status);
assert(driver_status == cudaDriverEntryPointSuccess);
return reinterpret_cast<PFN_cuTensorMapEncodeTiled_v12000>(cuTensorMapEncodeTiled_ptr);
}
template<int M_SMEM_TD, int N_SMEM_TD>
void tma_2d_map_c_init(PFN_cuTensorMapEncodeTiled_v12000 cuTensorMapEncodeTiled, CUtensorMap* tmap, void* ptr, uint64_t m_dim_gmem, uint64_t n_dim_gmem) {
constexpr uint32_t rank = 2;
uint64_t dim_gmem[rank] = {n_dim_gmem, m_dim_gmem};
uint64_t stride_gmem[rank - 1] = {n_dim_gmem * sizeof(__half)};
uint32_t dim_smem[rank] = {N_SMEM_TD, M_SMEM_TD};
uint32_t elem_stride[rank] = {1, 1};
// Create the tensor descriptor.
auto res = cuTensorMapEncodeTiled(
tmap, // CUtensorMap *tensorMap,
CUtensorMapDataType::CU_TENSOR_MAP_DATA_TYPE_FLOAT16,
rank, // cuuint32_t tensorRank,
ptr, // void *globalAddress,
dim_gmem, // const cuuint64_t *globalDim,
stride_gmem, // const cuuint64_t *globalStrides,
dim_smem, // const cuuint32_t *boxDim,
elem_stride, // const cuuint32_t *elementStrides,
CUtensorMapInterleave::CU_TENSOR_MAP_INTERLEAVE_NONE, // Interleave patterns can be used to accelerate loading of values that are less than 4 bytes long.
CUtensorMapSwizzle::CU_TENSOR_MAP_SWIZZLE_NONE, // Swizzling can be used to avoid shared memory bank conflicts.
CUtensorMapL2promotion::CU_TENSOR_MAP_L2_PROMOTION_NONE, // L2 Promotion can be used to widen the effect of a cache-policy to a wider set of L2 cache lines.
CUtensorMapFloatOOBfill::CU_TENSOR_MAP_FLOAT_OOB_FILL_NONE // Any element that is outside of bounds will be set to zero by the TMA transfer.
);
}
template<int MN_SMEM_TD, int K_SMEM_TD, CUtensorMapSwizzle SWIZZLE>
struct tma_3d_map_ab {
static void init(PFN_cuTensorMapEncodeTiled_v12000 cuTensorMapEncodeTiled, CUtensorMap* tmap, void* ptr, uint64_t mn_dim_gmem, uint64_t k_dim_gmem);
};
// For No swizzle canonical layout of 128b segments is 8x1 (or core matrices of 8 rows x 32 element columns (since 32 NVFP4 = 16B = 128b))
template<int MN_SMEM_TD, int K_SMEM_TD>
struct tma_3d_map_ab<MN_SMEM_TD, K_SMEM_TD, CUtensorMapSwizzle::CU_TENSOR_MAP_SWIZZLE_NONE> {
static void init(PFN_cuTensorMapEncodeTiled_v12000 cuTensorMapEncodeTiled, CUtensorMap* tmap, void* ptr, uint64_t mn_dim_gmem, uint64_t k_dim_gmem) {
constexpr uint32_t rank = 3;
uint64_t dim_gmem[rank] = {32, mn_dim_gmem, k_dim_gmem/32};
uint64_t stride_gmem[rank - 1] = {k_dim_gmem/2, 16};
uint32_t dim_smem[rank] = {32, MN_SMEM_TD, K_SMEM_TD/32};
uint32_t elem_stride[rank] = {1, 1, 1};
// Create the tensor descriptor.
auto res = cuTensorMapEncodeTiled(
tmap, // CUtensorMap *tensorMap,
CUtensorMapDataType::CU_TENSOR_MAP_DATA_TYPE_16U4_ALIGN8B,
rank, // cuuint32_t tensorRank,
ptr, // void *globalAddress,
dim_gmem, // const cuuint64_t *globalDim,
stride_gmem, // const cuuint64_t *globalStrides,
dim_smem, // const cuuint32_t *boxDim,
elem_stride, // const cuuint32_t *elementStrides,
CUtensorMapInterleave::CU_TENSOR_MAP_INTERLEAVE_NONE, // Interleave patterns can be used to accelerate loading of values that are less than 4 bytes long.
CUtensorMapSwizzle::CU_TENSOR_MAP_SWIZZLE_NONE, // Swizzling can be used to avoid shared memory bank conflicts.
CUtensorMapL2promotion::CU_TENSOR_MAP_L2_PROMOTION_NONE, // L2 Promotion can be used to widen the effect of a cache-policy to a wider set of L2 cache lines.
CUtensorMapFloatOOBfill::CU_TENSOR_MAP_FLOAT_OOB_FILL_NONE // Any element that is outside of bounds will be set to zero by the TMA transfer.
);
// ISSUE: Insert error check here on res
}
};
// For 128B swizzle canonical layout of 128b segments is 8x8 (or core matrices of 8 rows x 256 element columns)
template<int MN_SMEM_TD, int K_SMEM_TD>
struct tma_3d_map_ab<MN_SMEM_TD, K_SMEM_TD, CUtensorMapSwizzle::CU_TENSOR_MAP_SWIZZLE_128B> {
static void init(PFN_cuTensorMapEncodeTiled_v12000 cuTensorMapEncodeTiled, CUtensorMap* tmap, void* ptr, uint64_t mn_dim_gmem, uint64_t k_dim_gmem) {
constexpr uint32_t rank = 3;
uint64_t dim_gmem[rank] = {256, mn_dim_gmem, k_dim_gmem/256};
uint64_t stride_gmem[rank - 1] = {k_dim_gmem/2, 128};
uint32_t dim_smem[rank] = {256, MN_SMEM_TD, K_SMEM_TD/256};
uint32_t elem_stride[rank] = {1, 1, 1};
// Create the tensor descriptor.
auto res = cuTensorMapEncodeTiled(
tmap, // CUtensorMap *tensorMap,
CUtensorMapDataType::CU_TENSOR_MAP_DATA_TYPE_16U4_ALIGN8B,
rank, // cuuint32_t tensorRank,
ptr, // void *globalAddress,
dim_gmem, // const cuuint64_t *globalDim,
stride_gmem, // const cuuint64_t *globalStrides,
dim_smem, // const cuuint32_t *boxDim,
elem_stride, // const cuuint32_t *elementStrides,
CUtensorMapInterleave::CU_TENSOR_MAP_INTERLEAVE_NONE, // Interleave patterns can be used to accelerate loading of values that are less than 4 bytes long.
CUtensorMapSwizzle::CU_TENSOR_MAP_SWIZZLE_128B, // Swizzling can be used to avoid shared memory bank conflicts.
CUtensorMapL2promotion::CU_TENSOR_MAP_L2_PROMOTION_NONE, // L2 Promotion can be used to widen the effect of a cache-policy to a wider set of L2 cache lines.
CUtensorMapFloatOOBfill::CU_TENSOR_MAP_FLOAT_OOB_FILL_NONE // Any element that is outside of bounds will be set to zero by the TMA transfer.
);
// ISSUE: Insert error check here on res
}
};
// For 64B swizzle canonical layout of 128b segments is 8x4 (or core matrices of 8 rows x 128 element columns)
template<int MN_SMEM_TD, int K_SMEM_TD>
struct tma_3d_map_ab<MN_SMEM_TD, K_SMEM_TD, CUtensorMapSwizzle::CU_TENSOR_MAP_SWIZZLE_64B> {
static void init(PFN_cuTensorMapEncodeTiled_v12000 cuTensorMapEncodeTiled, CUtensorMap* tmap, void* ptr, uint64_t mn_dim_gmem, uint64_t k_dim_gmem) {
constexpr uint32_t rank = 3;
uint64_t dim_gmem[rank] = {128, mn_dim_gmem, k_dim_gmem/128};
uint64_t stride_gmem[rank - 1] = {k_dim_gmem/2, 64};
uint32_t dim_smem[rank] = {128, MN_SMEM_TD, K_SMEM_TD/128};
uint32_t elem_stride[rank] = {1, 1, 1};
// Create the tensor descriptor.
auto res = cuTensorMapEncodeTiled(
tmap, // CUtensorMap *tensorMap,
CUtensorMapDataType::CU_TENSOR_MAP_DATA_TYPE_16U4_ALIGN8B,
rank, // cuuint32_t tensorRank,
ptr, // void *globalAddress,
dim_gmem, // const cuuint64_t *globalDim,
stride_gmem, // const cuuint64_t *globalStrides,
dim_smem, // const cuuint32_t *boxDim,
elem_stride, // const cuuint32_t *elementStrides,
CUtensorMapInterleave::CU_TENSOR_MAP_INTERLEAVE_NONE, // Interleave patterns can be used to accelerate loading of values that are less than 4 bytes long.
CUtensorMapSwizzle::CU_TENSOR_MAP_SWIZZLE_64B, // Swizzling can be used to avoid shared memory bank conflicts.
CUtensorMapL2promotion::CU_TENSOR_MAP_L2_PROMOTION_NONE, // L2 Promotion can be used to widen the effect of a cache-policy to a wider set of L2 cache lines.
CUtensorMapFloatOOBfill::CU_TENSOR_MAP_FLOAT_OOB_FILL_NONE // Any element that is outside of bounds will be set to zero by the TMA transfer.
);
// ISSUE: Insert error check here on res
}
};
// For 32B swizzle canonical layout of 128b segments is 8x2 (or core matrices of 8 rows x 64 element columns)
template<int MN_SMEM_TD, int K_SMEM_TD>
struct tma_3d_map_ab<MN_SMEM_TD, K_SMEM_TD, CUtensorMapSwizzle::CU_TENSOR_MAP_SWIZZLE_32B> {
static void init(PFN_cuTensorMapEncodeTiled_v12000 cuTensorMapEncodeTiled, CUtensorMap* tmap, void* ptr, uint64_t mn_dim_gmem, uint64_t k_dim_gmem) {
constexpr uint32_t rank = 3;
uint64_t dim_gmem[rank] = {64, mn_dim_gmem, k_dim_gmem/64};
uint64_t stride_gmem[rank - 1] = {k_dim_gmem/2, 32};
uint32_t dim_smem[rank] = {64, MN_SMEM_TD, K_SMEM_TD/64};
uint32_t elem_stride[rank] = {1, 1, 1};
// Create the tensor descriptor.
auto res = cuTensorMapEncodeTiled(
tmap, // CUtensorMap *tensorMap,
CUtensorMapDataType::CU_TENSOR_MAP_DATA_TYPE_16U4_ALIGN8B,
rank, // cuuint32_t tensorRank,
ptr, // void *globalAddress,
dim_gmem, // const cuuint64_t *globalDim,
stride_gmem, // const cuuint64_t *globalStrides,
dim_smem, // const cuuint32_t *boxDim,
elem_stride, // const cuuint32_t *elementStrides,
CUtensorMapInterleave::CU_TENSOR_MAP_INTERLEAVE_NONE, // Interleave patterns can be used to accelerate loading of values that are less than 4 bytes long.
CUtensorMapSwizzle::CU_TENSOR_MAP_SWIZZLE_32B, // Swizzling can be used to avoid shared memory bank conflicts.
CUtensorMapL2promotion::CU_TENSOR_MAP_L2_PROMOTION_NONE, // L2 Promotion can be used to widen the effect of a cache-policy to a wider set of L2 cache lines.
CUtensorMapFloatOOBfill::CU_TENSOR_MAP_FLOAT_OOB_FILL_NONE // Any element that is outside of bounds will be set to zero by the TMA transfer.
);
// ISSUE: Insert error check here on res
}
};
template<int K_SMEM_TD>
void tma_3d_map_sf(PFN_cuTensorMapEncodeTiled_v12000 cuTensorMapEncodeTiled, CUtensorMap* tmap, void* ptr, uint64_t mn_dim_gmem, uint64_t k_dim_gmem) {
constexpr uint32_t rank = 3;
constexpr int K_SMEM_TD_SF = K_SMEM_TD / 16;
const int k_dim_gmem_sf = k_dim_gmem / 16;
uint64_t dim_gmem[rank] = {256, 2ULL * (k_dim_gmem_sf / 4), mn_dim_gmem / 128};
uint64_t stride_gmem[rank - 1] = {256, 512ULL * (k_dim_gmem_sf / 4)};
uint32_t dim_smem[rank] = {256, 2ULL * (K_SMEM_TD_SF / 4), 1};
uint32_t elem_stride[rank] = {1, 1, 1};
// Create the tensor descriptor.
auto res = cuTensorMapEncodeTiled(
tmap, // CUtensorMap *tensorMap,
CUtensorMapDataType::CU_TENSOR_MAP_DATA_TYPE_UINT8,
rank, // cuuint32_t tensorRank,
ptr, // void *globalAddress,
dim_gmem, // const cuuint64_t *globalDim,
stride_gmem, // const cuuint64_t *globalStrides,
dim_smem, // const cuuint32_t *boxDim,
elem_stride, // const cuuint32_t *elementStrides,
CUtensorMapInterleave::CU_TENSOR_MAP_INTERLEAVE_NONE, // Interleave patterns can be used to accelerate loading of values that are less than 4 bytes long.
CUtensorMapSwizzle::CU_TENSOR_MAP_SWIZZLE_NONE, // Swizzling can be used to avoid shared memory bank conflicts.
CUtensorMapL2promotion::CU_TENSOR_MAP_L2_PROMOTION_NONE, // L2 Promotion can be used to widen the effect of a cache-policy to a wider set of L2 cache lines.
CUtensorMapFloatOOBfill::CU_TENSOR_MAP_FLOAT_OOB_FILL_NONE // Any element that is outside of bounds will be set to zero by the TMA transfer.
);
if (res != CUDA_SUCCESS) {
printf("TMA encode failed with error %d\\n", res);
}
}
/*
For each warp:
LANES: How many rows of TMEM are loaded
WIDTH: How many bits in each row
REPT: WIDTH repeated REPT times
*/
template<int LANES, int WIDTH, int REPT>
__device__ void inline tcgen05_ld(float* regs, int tmem_addr);
// |
// V
// Specializations
template<>
__device__ void inline tcgen05_ld<32, 32, 16>(float* regs, int tmem_addr) {
asm volatile (
"tcgen05.ld.sync.aligned.32x32b.x16.b32 { %0, %1, %2, %3, %4, %5, %6, %7, "
" %8, %9, %10, %11, %12, %13, %14, %15}, [%16];"
: "=f"(regs[0]), "=f"(regs[1]), "=f"(regs[2]), "=f"(regs[3]),
"=f"(regs[4]), "=f"(regs[5]), "=f"(regs[6]), "=f"(regs[7]),
"=f"(regs[8]), "=f"(regs[9]), "=f"(regs[10]), "=f"(regs[11]),
"=f"(regs[12]), "=f"(regs[13]), "=f"(regs[14]), "=f"(regs[15])
: "r"(tmem_addr)
);
}
template<>
__device__ void inline tcgen05_ld<16, 256, 2>(float* regs, int tmem_addr) {
asm volatile (
"tcgen05.ld.sync.aligned.16x256b.x2.b32 { %0, %1, %2, %3, %4, %5, %6, %7 }, [%8]; "
: "=f"(regs[0]), "=f"(regs[1]), "=f"(regs[2]), "=f"(regs[3]),
"=f"(regs[4]), "=f"(regs[5]), "=f"(regs[6]), "=f"(regs[7])
: "r"(tmem_addr)
);
}
template<>
__device__ void inline tcgen05_ld<16, 256, 4>(float* regs, int tmem_addr) {
asm volatile (
"tcgen05.ld.sync.aligned.16x256b.x4.b32 { %0, %1, %2, %3, %4, %5, %6, %7, "
" %8, %9, %10, %11, %12, %13, %14, %15}, [%16]; "
: "=f"(regs[0]), "=f"(regs[1]), "=f"(regs[2]), "=f"(regs[3]),
"=f"(regs[4]), "=f"(regs[5]), "=f"(regs[6]), "=f"(regs[7]),
"=f"(regs[8]), "=f"(regs[9]), "=f"(regs[10]), "=f"(regs[11]),
"=f"(regs[12]), "=f"(regs[13]), "=f"(regs[14]), "=f"(regs[15])
: "r"(tmem_addr)
);
}
template<>
__device__ void inline tcgen05_ld<16, 256, 8>(float* regs, int tmem_addr) {
asm volatile (
"tcgen05.ld.sync.aligned.16x256b.x8.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"(regs[0]), "=f"(regs[1]), "=f"(regs[2]), "=f"(regs[3]),
"=f"(regs[4]), "=f"(regs[5]), "=f"(regs[6]), "=f"(regs[7]),
"=f"(regs[8]), "=f"(regs[9]), "=f"(regs[10]), "=f"(regs[11]),
"=f"(regs[12]), "=f"(regs[13]), "=f"(regs[14]), "=f"(regs[15]),
"=f"(regs[16]), "=f"(regs[17]), "=f"(regs[18]), "=f"(regs[19]),
"=f"(regs[20]), "=f"(regs[21]), "=f"(regs[22]), "=f"(regs[23]),
"=f"(regs[24]), "=f"(regs[25]), "=f"(regs[26]), "=f"(regs[27]),
"=f"(regs[28]), "=f"(regs[29]), "=f"(regs[30]), "=f"(regs[31])
: "r"(tmem_addr)
);
}
template<>
__device__ void inline tcgen05_ld<16, 256, 16>(float* regs, int tmem_addr) {
asm volatile (
"tcgen05.ld.sync.aligned.16x256b.x16.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"(regs[0]), "=f"(regs[1]), "=f"(regs[2]), "=f"(regs[3]),
"=f"(regs[4]), "=f"(regs[5]), "=f"(regs[6]), "=f"(regs[7]),
"=f"(regs[8]), "=f"(regs[9]), "=f"(regs[10]), "=f"(regs[11]),
"=f"(regs[12]), "=f"(regs[13]), "=f"(regs[14]), "=f"(regs[15]),
"=f"(regs[16]), "=f"(regs[17]), "=f"(regs[18]), "=f"(regs[19]),
"=f"(regs[20]), "=f"(regs[21]), "=f"(regs[22]), "=f"(regs[23]),
"=f"(regs[24]), "=f"(regs[25]), "=f"(regs[26]), "=f"(regs[27]),
"=f"(regs[28]), "=f"(regs[29]), "=f"(regs[30]), "=f"(regs[31]),
"=f"(regs[32]), "=f"(regs[33]), "=f"(regs[34]), "=f"(regs[35]),
"=f"(regs[36]), "=f"(regs[37]), "=f"(regs[38]), "=f"(regs[39]),
"=f"(regs[40]), "=f"(regs[41]), "=f"(regs[42]), "=f"(regs[43]),
"=f"(regs[44]), "=f"(regs[45]), "=f"(regs[46]), "=f"(regs[47]),
"=f"(regs[48]), "=f"(regs[49]), "=f"(regs[50]), "=f"(regs[51]),
"=f"(regs[52]), "=f"(regs[53]), "=f"(regs[54]), "=f"(regs[55]),
"=f"(regs[56]), "=f"(regs[57]), "=f"(regs[58]), "=f"(regs[59]),
"=f"(regs[60]), "=f"(regs[61]), "=f"(regs[62]), "=f"(regs[63])
: "r"(tmem_addr)
);
}
template<>
__device__ void inline tcgen05_ld<16, 256, 32>(float* regs, int tmem_addr) {
asm volatile (
"tcgen05.ld.sync.aligned.16x256b.x32.b32 { %0, %1, %2, %3, %4, %5, %6, %7, "
" %8, %9, %10, %11, %12, %13, %14, %15, "
" %16, %17, %18, %19, %20, %21, %22, %23, "
" %24, %25, %26, %27, %28, %29, %30, %31, "
" %32, %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"(regs[0]), "=f"(regs[1]), "=f"(regs[2]), "=f"(regs[3]),
"=f"(regs[4]), "=f"(regs[5]), "=f"(regs[6]), "=f"(regs[7]),
"=f"(regs[8]), "=f"(regs[9]), "=f"(regs[10]), "=f"(regs[11]),
"=f"(regs[12]), "=f"(regs[13]), "=f"(regs[14]), "=f"(regs[15]),
"=f"(regs[16]), "=f"(regs[17]), "=f"(regs[18]), "=f"(regs[19]),
"=f"(regs[20]), "=f"(regs[21]), "=f"(regs[22]), "=f"(regs[23]),
"=f"(regs[24]), "=f"(regs[25]), "=f"(regs[26]), "=f"(regs[27]),
"=f"(regs[28]), "=f"(regs[29]), "=f"(regs[30]), "=f"(regs[31]),
"=f"(regs[32]), "=f"(regs[33]), "=f"(regs[34]), "=f"(regs[35]),
"=f"(regs[36]), "=f"(regs[37]), "=f"(regs[38]), "=f"(regs[39]),
"=f"(regs[40]), "=f"(regs[41]), "=f"(regs[42]), "=f"(regs[43]),
"=f"(regs[44]), "=f"(regs[45]), "=f"(regs[46]), "=f"(regs[47]),
"=f"(regs[48]), "=f"(regs[49]), "=f"(regs[50]), "=f"(regs[51]),
"=f"(regs[52]), "=f"(regs[53]), "=f"(regs[54]), "=f"(regs[55]),
"=f"(regs[56]), "=f"(regs[57]), "=f"(regs[58]), "=f"(regs[59]),
"=f"(regs[60]), "=f"(regs[61]), "=f"(regs[62]), "=f"(regs[63]),
"=f"(regs[64]), "=f"(regs[65]), "=f"(regs[66]), "=f"(regs[67]),
"=f"(regs[68]), "=f"(regs[69]), "=f"(regs[70]), "=f"(regs[71]),
"=f"(regs[72]), "=f"(regs[73]), "=f"(regs[74]), "=f"(regs[75]),
"=f"(regs[76]), "=f"(regs[77]), "=f"(regs[78]), "=f"(regs[79]),
"=f"(regs[80]), "=f"(regs[81]), "=f"(regs[82]), "=f"(regs[83]),
"=f"(regs[84]), "=f"(regs[85]), "=f"(regs[86]), "=f"(regs[87]),
"=f"(regs[88]), "=f"(regs[89]), "=f"(regs[90]), "=f"(regs[91]),
"=f"(regs[92]), "=f"(regs[93]), "=f"(regs[94]), "=f"(regs[95]),
"=f"(regs[96]), "=f"(regs[97]), "=f"(regs[98]), "=f"(regs[99]),
"=f"(regs[100]), "=f"(regs[101]), "=f"(regs[102]), "=f"(regs[103]),
"=f"(regs[104]), "=f"(regs[105]), "=f"(regs[106]), "=f"(regs[107]),
"=f"(regs[108]), "=f"(regs[109]), "=f"(regs[110]), "=f"(regs[111]),
"=f"(regs[112]), "=f"(regs[113]), "=f"(regs[114]), "=f"(regs[115]),
"=f"(regs[116]), "=f"(regs[117]), "=f"(regs[118]), "=f"(regs[119]),
"=f"(regs[120]), "=f"(regs[121]), "=f"(regs[122]), "=f"(regs[123]),
"=f"(regs[124]), "=f"(regs[125]), "=f"(regs[126]), "=f"(regs[127])
: "r"(tmem_addr)
);
}
// Copies 32 rows x 128 bits from matrix described in SMEM by desc
// into tmem_ptr
template<int CTA_GROUP>
__device__ void inline tcgen05_cp(int tmem_ptr, uint64_t desc) {
asm volatile (
"tcgen05.cp.cta_group::%2.32x128b.warpx4 [%0], %1;"
:
: "r"(tmem_ptr), "l"(desc), "n"(CTA_GROUP)
);
}
__device__ uint64_t inline constexpr encode(uint64_t x) {
return (x & 0x3FFFF) >> 4;
}
/*
Un-changing instruction descriptor for tcgen05.mma
[0-1] : 0 (reserved)
[2] : Sparsity -> Dense = 0
[3] : 0 (reserved)
[4-5] : Matrix B Scale Factor Data ID -> always 0 for 4 SFs (using all bytes in each TMEM col)
[6] : 0 (reserved)
[7-9] : atype (Matrix A type) -> E2M1 = 1
[10-11] : btype (Matrix B type) -> E2M1 = 1
[12] : 0 (reserved)
[13] : Negate A Matrix -> 0 (no negation)
[14] : Negate B Matrix -> 0 (no negation)
[15] : Transpose A Matrix -> 0 (transposition not allowed, nor wanted)
[16] : Transpose B Matrix -> 0 (^^^^)
[17-22] : N, Dimension of Matrix B (3 LSBs not included) -> N >> 3
[23] : Scale Matrix Type, for both scale_A / scale_B -> UE4M3 = 0
[24-26] : 0 (reserved)
[27-28] : M, Dimension of Matrix A (7 LSBs not included) -> M >> 7
[29-30] : Matrix A Scale Factor Data ID -> always 0 (same as above)
[31] : K Dimension -> 0 with Dense from bit 2 makes desired K=64
*/
template<int M, int N>
__device__ uint32_t constexpr make_instr_desc() {
return (1 << 7) | (1 << 10) | ((N >> 3) << 17) | ((M >> 7) << 27);
}
// Complete descriptor with address info
template<int MN_DIM, CUtensorMapSwizzle SWIZZLE_TYPE = CUtensorMapSwizzle::CU_TENSOR_MAP_SWIZZLE_NONE>
__device__ uint64_t inline make_smem_desc(int smem_addr) {
constexpr uint64_t LBO = SWIZZLE_TYPE != CUtensorMapSwizzle::CU_TENSOR_MAP_SWIZZLE_NONE ? 1 : MN_DIM*16;
constexpr uint64_t SBO = SWIZZLE_TYPE == CUtensorMapSwizzle::CU_TENSOR_MAP_SWIZZLE_128B ? 8 * 128 :
(SWIZZLE_TYPE == CUtensorMapSwizzle::CU_TENSOR_MAP_SWIZZLE_64B ? 8 * 64 :
(SWIZZLE_TYPE == CUtensorMapSwizzle::CU_TENSOR_MAP_SWIZZLE_32B ? 8 * 32 : 8 * 16));
constexpr uint64_t SWIZZLE_BITS = SWIZZLE_TYPE == CUtensorMapSwizzle::CU_TENSOR_MAP_SWIZZLE_128B ? 2 :
(SWIZZLE_TYPE == CUtensorMapSwizzle::CU_TENSOR_MAP_SWIZZLE_64B ? 4 :
(SWIZZLE_TYPE == CUtensorMapSwizzle::CU_TENSOR_MAP_SWIZZLE_32B ? 6 : 0));
return encode(smem_addr) | (encode(LBO) << 16) | (encode(SBO) << 32) | (0x1ULL << 46) | (SWIZZLE_BITS << 61);
}
template<int CTA_GROUP>
__device__ void inline tcgen05_dealloc_tmem(int tmem_addr, int n_cols) {
asm volatile(
"tcgen05.dealloc.cta_group::%2.sync.aligned.b32 %0, %1;"
:
: "r"(tmem_addr), "r"(n_cols), "n"(CTA_GROUP)
);
}
// Warp synchronous execution (all threads in a warp execute)
template<int CTA_GROUP>
__device__ void inline tcgen05_alloc_tmem(int *tmem_addr_ptr, const int n_cols) {
// Performs a cvt.u64.u32, enables proper passing of smem ptr to PTX assembly
const int tmem_addr_ptr_cvt = static_cast<int>(__cvta_generic_to_shared(tmem_addr_ptr));
asm volatile (
"tcgen05.alloc.cta_group::%2.sync.aligned.shared::cta.b32 [%0], %1;"
:
: "r"(tmem_addr_ptr_cvt), "r"(n_cols), "n"(CTA_GROUP)
);
}
// Single thread execution
template<int CTA_GROUP>
__device__ void inline tcgen05_mma_nvfp4(int d_tmem, uint64_t a_desc, uint64_t b_desc, uint32_t i_desc,
int sfa_tmem, int sfb_tmem, int enable_input_d) {
asm volatile (
"{\\n"
".reg .pred p;\\n"
"setp.ne.b32 p, %4, 0;\\n"
"tcgen05.mma.cta_group::%7.kind::mxf4nvf4.block_scale.block16 [%0], %1, %2, %3, [%5], [%6], p; \\n"
"}\\n"
:
: "r"(d_tmem), "l"(a_desc), "l"(b_desc), "r"(i_desc), "r"(enable_input_d),
"r"(sfa_tmem), "r"(sfb_tmem), "n"(CTA_GROUP)
);
}
__device__ inline int clc_check_steal(const int clc_result_addr) {
int res = -1;
asm volatile(
"{\\n"
".reg .pred P1;\\n"
".reg .b128 handle;\\n"
"ld.shared.b128 handle, [%1];\\n"
"clusterlaunchcontrol.query_cancel.is_canceled.pred.b128 P1, handle;\\n"
"@!P1 bra.uni DONE;\\n" // if query returned false, no more work tiles to be stolen, so return -1 by doing nothing
"clusterlaunchcontrol.query_cancel.get_first_ctaid::x.b32.b128 %0, handle;\\n" // set res to stolen blockIdx.x
"DONE:\\n"
"}"
: "+r"(res)
: "r"(clc_result_addr)
);
return res;
}
__device__ inline void tmap_update_addr(const int local_tmap_addr, const void* new_addr) {
// Update base addresses
asm volatile(
"tensormap.replace.tile.global_address.shared::cta.b1024.b64 [%0], %1;"
:
: "r"(local_tmap_addr), "l"(new_addr)
);
}
template<int DIM>
__device__ inline void tmap_update_dim(const int local_tmap_addr, const int new_dim) {
// Adjust DIM value
asm volatile(
"tensormap.replace.tile.global_dim.shared::cta.b1024.b32 [%0], %1, %2;"
:
: "r"(local_tmap_addr), "n"(DIM), "r"(new_dim)
);
}
template<int DIM = 1>
__device__ inline void tmap_update_stride(const int local_tmap_addr, const uint64_t new_stride) {
// Adjust stride
asm volatile(
"tensormap.replace.tile.global_stride.shared::cta.b1024.b64 [%0], %1, %2;"
:
: "r"(local_tmap_addr), "n"(DIM), "l"(new_stride)
);
}
__device__ inline void tmap_update(const int local_tmap_addr, const void* new_addr) {
// Update base addresses
asm volatile(
"tensormap.replace.tile.global_address.shared::cta.b1024.b64 [%0], %1;"
:
: "r"(local_tmap_addr), "l"(new_addr)
);
}
template<int DIM = 1>
__device__ inline void tmap_update(const int local_tmap_addr, const void* new_addr, const int new_dim) {
// Adjust M-dim value
asm volatile(
"tensormap.replace.tile.global_dim.shared::cta.b1024.b32 [%0], %1, %2;"
:
: "r"(local_tmap_addr), "n"(DIM), "r"(new_dim)
);
// Update base addresses
asm volatile(
"tensormap.replace.tile.global_address.shared::cta.b1024.b64 [%0], %1;"
:
: "r"(local_tmap_addr), "l"(new_addr)
);
}
template<int DIM = 1>
__device__ inline void tmap_update(const int local_tmap_addr, const void* new_addr, const int new_dim, const uint64_t new_stride) {
// Adjust M-dim value
asm volatile(
"tensormap.replace.tile.global_dim.shared::cta.b1024.b32 [%0], %1, %2;"
:
: "r"(local_tmap_addr), "n"(DIM), "r"(new_dim)
);
// Adjust stride
asm volatile(
"tensormap.replace.tile.global_stride.shared::cta.b1024.b64 [%0], %1, %2;"
:
: "r"(local_tmap_addr), "n"(0), "l"(new_stride)
);
// Update base addresses
asm volatile(
"tensormap.replace.tile.global_address.shared::cta.b1024.b64 [%0], %1;"
:
: "r"(local_tmap_addr), "l"(new_addr)
);
}
__device__ inline void tmap_fence_proxy(const CUtensorMap* g_tmap, const int local_tmap_addr) {
asm volatile(
"tensormap.cp_fenceproxy.global.shared::cta.tensormap::generic.release.gpu.sync.aligned [%0], [%1], 128;"
:
: "l"(g_tmap), "r"(local_tmap_addr)
);
asm volatile(
"fence.proxy.tensormap::generic.acquire.gpu [%0], 128;"
:
: "l"(g_tmap)
);
}
template<int A, int B>
int constexpr MAX() {
if (A < B)
return B;
return A;
}
struct GroupDesc {
void* A_addr;
void* B_addr;
__half* C_addr;
uint8_t* sfa_addr;
uint8_t* sfb_addr;
int M;
int N;
int K;
int block_start;
};
constexpr int MAX_G = 8;
struct GroupDescs {
GroupDesc groups[MAX_G];
};
constexpr int WARP_SIZE = 32;
constexpr int SF_BLOCK_SIZE = 16;
#define DEBUG 0
/*
Warp 0 will be responsible for all single thread (async) issued instructions, warp 1 or all CTA threads will handle the rest
We need 1 warp to do the computation and TMA transfers. We need 4 warps in order to read/write all of TMEM (a single warp can only
access 32 lanes (rows) in TMEM out of the total 128 per SM)
For now we assume TD_CTA_M/N == TD_SMEM_M/N so each CTA computes just one output tile
Also we assume TD_SMEM_M/N == TD_MMA_M/N to avoid excess copies from SMEM back to TMEM, although implementing the modifications
To allow differing sizes shouldn't be too hard
ISSUE: NUM_WARPS should match launch conditions (is there a cleaner way to handle this, maybe launch bounds?)
*/
template<int TD_CTA_M, int TD_CTA_N,
int TD_SMEM_M, int TD_SMEM_N, int TD_SMEM_K,
int TD_MMA_M, int TD_MMA_N, int TD_MMA_K, CUtensorMapSwizzle SWIZZLE_TYPE,
int PIPE_STAGES, int NUM_WARPS, int OUT_N_CHUNK, int TILE_DESC_PIPE_STAGES,
int CLUSTER_SIZE, bool SINGLE_WAVE, bool NK_VAR>
__global__ void __cluster_dims__(CLUSTER_SIZE, 1, 1) nvfp4_group_gemm_kernel(const __grid_constant__ GroupDescs group_descs, const __grid_constant__ CUtensorMap tmap_a_temp,
const __grid_constant__ CUtensorMap tmap_b_temp,
const __grid_constant__ CUtensorMap tmap_c_temp,
const int total_tiles, CUtensorMap* d_tmaps,
int N, int K, const int G) {
const GroupDesc* groups = group_descs.groups;
// Statically computed values
constexpr int WIDTH_COREMAT = SWIZZLE_TYPE == CUtensorMapSwizzle::CU_TENSOR_MAP_SWIZZLE_128B ? 256 :
(SWIZZLE_TYPE == CUtensorMapSwizzle::CU_TENSOR_MAP_SWIZZLE_64B ? 128 :
(SWIZZLE_TYPE == CUtensorMapSwizzle::CU_TENSOR_MAP_SWIZZLE_32B ? 64 : 32));
constexpr int SF_SMEM_TILESZ = 128 * (TD_SMEM_K / SF_BLOCK_SIZE);
constexpr int A_SMEM_TILESZ = TD_SMEM_M * (TD_SMEM_K / 2);
constexpr int B_SMEM_TILESZ = TD_SMEM_N * (TD_SMEM_K / 2);
constexpr int SMEM_TILE_SZ = A_SMEM_TILESZ + B_SMEM_TILESZ + SF_SMEM_TILESZ + (1 + TD_MMA_N/256)*SF_SMEM_TILESZ;
constexpr int C_CHUNK_SMEM_TILESZ = TD_SMEM_M * OUT_N_CHUNK; // Size of chunk used for double buffer
//constexpr int C_SMEM_TILESZ = TD_SMEM_M * TD_SMEM_N;
constexpr uint16_t CLUSTER_MASK = ((uint16_t)1 << CLUSTER_SIZE) - 1; // Set first CLUSTER_SIZE bits to 1
// Calculate constants, offsets, etc... for this thread/warp
const int warp_id = threadIdx.x / WARP_SIZE;
const int lane_id = threadIdx.x % WARP_SIZE;
// Allocate SMEM buffers
/*
All Buffers need to be 128B aligned for TMA transfers
Need TD_SMEM_M * TD_SMEM_K * sizeof(nvfp4) bytes for a_smem
Need TD_SMEM_M * (TD_SMEM_K/SF_BLOCK_SIZE) * sizeof(fp8) bytes for sfa_smem
Need TD_SMEM_N * TD_SMEM_K * sizeof(nvfp4) bytes for b_smem
Need TD_SMEM_N * (TD_SMEM_K/SF_BLOCK_SIZE) * sizeof(fp8) bytes for sfb_smem
*/
__shared__ alignas(128) char a_smem[PIPE_STAGES * A_SMEM_TILESZ];
__shared__ alignas(128) char sfa_smem[PIPE_STAGES * SF_SMEM_TILESZ];
__shared__ alignas(128) char b_smem[PIPE_STAGES * B_SMEM_TILESZ];
__shared__ alignas(128) char sfb_smem[PIPE_STAGES * ((1 + TD_MMA_N/256) * SF_SMEM_TILESZ)];
__shared__ alignas(128) __half c_smem[C_CHUNK_SMEM_TILESZ * 2]; // *2 for double buffer output
// Convert ptrs to properly pass to inline PTX
const int a_smem_ptr = static_cast<int>(__cvta_generic_to_shared(a_smem));
const int sfa_smem_ptr = static_cast<int>(__cvta_generic_to_shared(sfa_smem));
const int b_smem_ptr = static_cast<int>(__cvta_generic_to_shared(b_smem));
const int sfb_smem_ptr = static_cast<int>(__cvta_generic_to_shared(sfb_smem));
const int c_smem_ptr = static_cast<int>(__cvta_generic_to_shared(c_smem));
// GMEM cache for CTA specific tensor maps
CUtensorMap* g_A_tmap = d_tmaps + 3 * TILE_DESC_PIPE_STAGES * blockIdx.x;
CUtensorMap* g_B_tmap = g_A_tmap + TILE_DESC_PIPE_STAGES;
CUtensorMap* g_C_tmap = g_B_tmap + TILE_DESC_PIPE_STAGES;
__shared__ int m_off_arr[TILE_DESC_PIPE_STAGES];
__shared__ int n_off_arr[TILE_DESC_PIPE_STAGES];
__shared__ alignas(128) CUtensorMap local_A_tmap; // ISSUE: Do these need to be aligned?
__shared__ alignas(128) CUtensorMap local_B_tmap;
__shared__ alignas(128) CUtensorMap local_C_tmap;
if (warp_id == 0 && elect_one_sync()) {
local_A_tmap = tmap_a_temp;
local_B_tmap = tmap_b_temp;
local_C_tmap = tmap_c_temp;
}
const int local_A_tmap_addr = static_cast<int>(__cvta_generic_to_shared(&local_A_tmap));
const int local_B_tmap_addr = static_cast<int>(__cvta_generic_to_shared(&local_B_tmap));
const int local_C_tmap_addr = static_cast<int>(__cvta_generic_to_shared(&local_C_tmap));
// Allocate TMEM buffers, single warp execution ISSUE: IMPLEMENT TWO ALLOC TECHNIQUE WITH ONE FOR RESULT ONE FOR SF TILES (THERE ARE TRADEOFFS WITH ALLOCATION SPACE VS NUM ALLOCATIONS, COULD BE A PROB SIZE SPECIFIC THING)
/*
We need TD_MMA_N columns for the result (using FP32 accumulation)
SFA needs TD_MMA_M/32 columns per 64 elements in K (due to (32x4)xcols layout discussed previously) -> (TD_MMA_M/32) * (TD_SMEM_K/64) total columns
SFB needs TD_MMA_N/32 columns per 64 elements in K (due to (32x4)xcols layout) -> (TD_MMA_N/32) * (TD_SMEM_K/64) total columns
TMEM Address Structure:
[0-15] : Column Index
[31-16] : Lane Index
MMA result buffer
*/
__shared__ int tmem_addr_ptr[1];
// Setup memory barriers
__shared__ alignas(8) int64_t mbars[PIPE_STAGES * 2 + 4 + TILE_DESC_PIPE_STAGES];
const int mbar_addr_tma = static_cast<int>(__cvta_generic_to_shared(mbars));
const int mbar_addr_mma = mbar_addr_tma + PIPE_STAGES * 8; // 8 because each mbar is 64bits = 8B
const int mbar_addr_epi = mbar_addr_mma + PIPE_STAGES * 8;
const int mbar_addr_epi_done = mbar_addr_epi + 2 * 8;
const int mbar_addr_tile_ready = mbar_addr_epi_done + 2 * 8;
int epi_phase[2] = {0, 0};
int epi_done_phase[2] = {0, 0};
int tile_ready_phase[TILE_DESC_PIPE_STAGES] = {};
int tmem_buf = 0;
if (warp_id == 0 && elect_one_sync()) {
for (int i = 0; i < PIPE_STAGES * 2 + 4 + TILE_DESC_PIPE_STAGES; i++) {
mbar_init(mbar_addr_tma + i * 8, (i >= PIPE_STAGES && i < 2 * PIPE_STAGES) ? CLUSTER_SIZE : 1);
}
asm volatile("fence.mbarrier_init.release.cluster;"); // ISSUE: Verify we need this here
} else if (warp_id == 1) {
tcgen05_alloc_tmem<1>(tmem_addr_ptr, TD_MMA_N*2*2);
}
__syncthreads(); // Ensure all threads have correct TMEM ptrs
const int tmem_addr_base = tmem_addr_ptr[0];
const int tmem_result_ptrs[2] = {tmem_addr_base, tmem_addr_base + 2*TD_MMA_N};
const int tmem_sfa_ptrs[2] = {tmem_result_ptrs[0] + TD_MMA_N, tmem_result_ptrs[1] + TD_MMA_N};
const int tmem_sfb_ptrs[2] = {tmem_sfa_ptrs[0] + (TD_MMA_M/32) * (TD_SMEM_K/64), tmem_sfa_ptrs[1] + (TD_MMA_M/32) * (TD_SMEM_K/64)};
constexpr int TILE_WARP = NUM_WARPS - 3;
constexpr int TMA_WARP = NUM_WARPS - 2;
constexpr int MMA_WARP = NUM_WARPS - 1;
// Work-tile loop
int n_tiles = CEIL_DIV(N, TD_CTA_N);
__shared__ uint8_t* sfa_gmem_base_arr[TILE_DESC_PIPE_STAGES];
__shared__ uint8_t* sfb_gmem_base_arr[TILE_DESC_PIPE_STAGES];
__shared__ int M_vals[TILE_DESC_PIPE_STAGES];
__shared__ int N_vals[NK_VAR ? TILE_DESC_PIPE_STAGES : 1];
__shared__ int K_vals[NK_VAR ? TILE_DESC_PIPE_STAGES : 1];
int glob_k_off = 0;
int tile_stages = 0;
int cta_rank;
asm volatile("mov.b32 %0, %%cluster_ctarank;" : "=r"(cta_rank));
const bool is_leader = (cta_rank == 0);
for (int tile_idx = blockIdx.x; tile_idx < total_tiles; tile_idx += gridDim.x) {
int tile_stage = tile_stages % TILE_DESC_PIPE_STAGES;
if (warp_id == TILE_WARP) {
// ISSUE: if stages wrap around we need another set of barriers to ensure consumers have consumed
int group, lo = 0, hi = G - 1;
while (lo <= hi) {
group = lo + ((hi - lo) / 2);
if (tile_idx < groups[group].block_start) {
hi = group - 1; // search left
} else if (group < G - 1 && tile_idx >= groups[group + 1].block_start) {
lo = group + 1; // search right
} else {
break; // found group of tile_idx;
}
}
int group_tile = tile_idx - groups[group].block_start;
if constexpr (NK_VAR) {
n_tiles = CEIL_DIV(groups[group].N, TD_CTA_N);
}
int row_idx = group_tile / n_tiles;
int col_idx = group_tile % n_tiles;
m_off_arr[tile_stage] = row_idx * TD_CTA_M;
n_off_arr[tile_stage] = col_idx * TD_CTA_N;
M_vals[tile_stage] = groups[group].M;
if constexpr (NK_VAR) {
N_vals[tile_stage] = groups[group].N;
K_vals[tile_stage] = groups[group].K;
}
sfa_gmem_base_arr[tile_stage] = groups[group].sfa_addr;
sfb_gmem_base_arr[tile_stage] = groups[group].sfb_addr;
asm volatile ("fence.proxy.async.shared::cta;" ::: "memory"); // Ensure SMEM writes are visible
if (elect_one_sync()) {
if constexpr (NK_VAR) {
tmap_update_addr(local_A_tmap_addr, groups[group].A_addr);
tmap_update_dim<1>(local_A_tmap_addr, groups[group].M);
tmap_update_dim<2>(local_A_tmap_addr, groups[group].K / WIDTH_COREMAT);
tmap_update_stride<0>(local_A_tmap_addr, groups[group].K / 2);
tmap_update_addr(local_B_tmap_addr, groups[group].B_addr);
tmap_update_dim<1>(local_B_tmap_addr, groups[group].N);
tmap_update_dim<2>(local_B_tmap_addr, groups[group].K / WIDTH_COREMAT);
tmap_update_stride<0>(local_B_tmap_addr, groups[group].K / 2);
tmap_update_addr(local_C_tmap_addr, groups[group].C_addr);
tmap_update_dim<0>(local_C_tmap_addr, groups[group].N);
tmap_update_dim<1>(local_C_tmap_addr, groups[group].M);
tmap_update_stride<0>(local_C_tmap_addr, groups[group].N * sizeof(__half));
} else {
tmap_update(local_A_tmap_addr, groups[group].A_addr, groups[group].M);
tmap_update(local_B_tmap_addr, groups[group].B_addr);
tmap_update(local_C_tmap_addr, groups[group].C_addr, groups[group].M);
}
}
__syncwarp();
tmap_fence_proxy(g_A_tmap + tile_stage, local_A_tmap_addr);
tmap_fence_proxy(g_B_tmap + tile_stage, local_B_tmap_addr);
tmap_fence_proxy(g_C_tmap + tile_stage, local_C_tmap_addr);
// Ensure all SMEM writes are visible and signal barrier
if (elect_one_sync()) {
mbar_arrive(mbar_addr_tile_ready + 8 * tile_stage, 1);
}
}
// All other warps wait for tile info to be ready
mbar_wait(mbar_addr_tile_ready + 8 * tile_stage, tile_ready_phase[tile_stage]);
tile_ready_phase[tile_stage] ^= 1;
int m_off = m_off_arr[tile_stage];
int n_off = n_off_arr[tile_stage];
uint8_t* sfa_gmem_base = sfa_gmem_base_arr[tile_stage];
uint8_t* sfb_gmem_base = sfb_gmem_base_arr[tile_stage];
if constexpr (NK_VAR) {
N = N_vals[tile_stage];
K = K_vals[tile_stage];
}
// TMA thread loops over SMEM tile stages and loads from GMEM->SMEM
if (warp_id == TMA_WARP) {
if (elect_one_sync()) {
auto tma_load_stage = [&](const int k_off, const int stage) {
const int k_off_coremat = k_off / WIDTH_COREMAT;
const int a_smem_stage_ptr = a_smem_ptr + stage * A_SMEM_TILESZ;
const int b_smem_stage_ptr = b_smem_ptr + stage * B_SMEM_TILESZ;
const int sfa_smem_stage_ptr = sfa_smem_ptr + stage * SF_SMEM_TILESZ;
const int sfb_smem_stage_ptr = sfb_smem_ptr + stage * SF_SMEM_TILESZ * (TD_MMA_N == 256 ? 2 : 1);
const int mbar_addr_tma_stage = mbar_addr_tma + stage * 8;
/*
Scale factors are stored in global memory in 4x4x32 chunks, i.e. 512B chunks where each chunk represents a
128x4 chunk of the SF matrix (in M or N xK)
So we calculate the offset in each dimension in terms of these 512B chunks:
k_off / 64 represents the number of 128x4 (512B) chunks along the K dimension which are contiguous (4 * SF_BLOCKS_SIZE = 64)
m/n_off / 128 represents the number of 512B chunks along the M dimension which are strided by K / 64 512B chunks
*/
const uint8_t* sfa_g_ptr = sfa_gmem_base + ((k_off / 64) + (m_off / 128) * (K / 64)) * 512; // ISSUE: These could be just simple bit shifts, adjust if compiler doesn't
const uint8_t* sfb_g_ptr = sfb_gmem_base + ((k_off / 64) + (n_off / 128) * (K / 64)) * 512;
// Signal that we expect SMEM_TILE_SZ bytes to arrive on the local mbar object before proceeding to next phase
mbar_arrive_expect(mbar_addr_tma_stage, SMEM_TILE_SZ);
if (is_leader) {
// Leader CTA issues multicast TMA
tcgen05_3dtma_g2s_ab_multicast<1, CLUSTER_MASK>(a_smem_stage_ptr, g_A_tmap + tile_stage, m_off, k_off_coremat, mbar_addr_tma_stage, CacheHintSm100::EVICT_NORMAL);
}
tcgen05_3dtma_g2s_ab<1>(b_smem_stage_ptr, g_B_tmap + tile_stage, n_off, k_off_coremat, mbar_addr_tma_stage, CacheHintSm100::EVICT_NORMAL);
tcgen05_1dtma_g2s_sf(sfa_smem_stage_ptr, sfa_g_ptr, SF_SMEM_TILESZ, mbar_addr_tma_stage, CacheHintSm100::EVICT_NORMAL);
tcgen05_1dtma_g2s_sf(sfb_smem_stage_ptr, sfb_g_ptr, SF_SMEM_TILESZ, mbar_addr_tma_stage, CacheHintSm100::EVICT_NORMAL);
if constexpr (TD_MMA_N == 256) {
tcgen05_1dtma_g2s_sf(sfb_smem_stage_ptr+SF_SMEM_TILESZ, sfb_g_ptr + (K / 64)*512, SF_SMEM_TILESZ, mbar_addr_tma_stage, CacheHintSm100::EVICT_NORMAL);
}
};
int k_off = 0;
// If first tile kick off pipeline without waiting on MMA
if (tile_idx == blockIdx.x) {
const int k_tiles = K / TD_SMEM_K;
const int prefetch_stages = PIPE_STAGES > k_tiles ? k_tiles : PIPE_STAGES;
for (int stage = 0; stage < prefetch_stages; stage++) {
tma_load_stage(stage * TD_SMEM_K, stage);
}
k_off = TD_SMEM_K * prefetch_stages;
}
// Cycle through tile stages, loading tiles once no longer in use by the MMA stage
glob_k_off += k_off;
for (; k_off < K; k_off += TD_SMEM_K) {
const int stage = (glob_k_off / TD_SMEM_K) % PIPE_STAGES;
mbar_wait(mbar_addr_mma + stage * 8, (((glob_k_off / TD_SMEM_K) / PIPE_STAGES) - 1) % 2);
tma_load_stage(k_off, stage);
glob_k_off += TD_SMEM_K;
}
}
}
// MMA thread loops over SMEM tile stages, loads TMEM and computes MMA ops
else if (warp_id == MMA_WARP && elect_one_sync()) {
const int tmem_addr_result = tmem_result_ptrs[tmem_buf];
const int tmem_addr_sfa = tmem_sfa_ptrs[tmem_buf];
const int tmem_addr_sfb = tmem_sfb_ptrs[tmem_buf];
for (int k_off = 0; k_off < K; k_off += TD_SMEM_K) {
const int stage = (glob_k_off / TD_SMEM_K) % PIPE_STAGES;
mbar_wait(mbar_addr_tma + stage * 8, ((glob_k_off / TD_SMEM_K) / PIPE_STAGES) % 2);
const int a_smem_stage_ptr = a_smem_ptr + stage * A_SMEM_TILESZ;
const int b_smem_stage_ptr = b_smem_ptr + stage * B_SMEM_TILESZ;
const int sfa_smem_stage_ptr = sfa_smem_ptr + stage * SF_SMEM_TILESZ;
const int sfb_smem_stage_ptr = sfb_smem_ptr + stage * SF_SMEM_TILESZ * (TD_MMA_N == 256 ? 2 : 1);
const int mbar_addr_mma_stage = mbar_addr_mma + stage * 8;
// Load scale factors SMEM -> TMEM
for (int sub_k_iter = 0; sub_k_iter < TD_SMEM_K / TD_MMA_K; sub_k_iter++) {
uint64_t sfa_desc = make_smem_desc<0>(sfa_smem_stage_ptr + (sub_k_iter * 512)); // ISSUE: verify this should input 0 here
uint64_t sfb_desc = make_smem_desc<0>(sfb_smem_stage_ptr + (sub_k_iter * 512));
tcgen05_cp<1>(tmem_addr_sfa + 4 * sub_k_iter, sfa_desc);
tcgen05_cp<1>(tmem_addr_sfb + MAX<TD_MMA_N / 32, 4>() * sub_k_iter, sfb_desc);
if constexpr (TD_MMA_N == 256) {
uint64_t sfb_desc2 = make_smem_desc<0>(sfb_smem_stage_ptr + SF_SMEM_TILESZ + (sub_k_iter * 512));
tcgen05_cp<1>(tmem_addr_sfb + 8 * sub_k_iter + 4, sfb_desc2);
}
}
// Loop over SMEM tile K-dim
for (int sub_k_iter = 0; sub_k_iter < TD_SMEM_K / TD_MMA_K; sub_k_iter++) {
// Stride computed differently depending on swizzle mode because it changes core matrix shape
uint64_t a_desc, b_desc;
if constexpr (SWIZZLE_TYPE != CUtensorMapSwizzle::CU_TENSOR_MAP_SWIZZLE_NONE) {
a_desc = make_smem_desc<TD_MMA_M, SWIZZLE_TYPE>(a_smem_stage_ptr + sub_k_iter * 32);
b_desc = make_smem_desc<TD_MMA_N, SWIZZLE_TYPE>(b_smem_stage_ptr + sub_k_iter * 32);
}
else {
a_desc = make_smem_desc<TD_MMA_M, SWIZZLE_TYPE>(a_smem_stage_ptr + sub_k_iter * TD_MMA_K * TD_MMA_M / 2);
b_desc = make_smem_desc<TD_MMA_N, SWIZZLE_TYPE>(b_smem_stage_ptr + sub_k_iter * TD_MMA_K * TD_MMA_N / 2);
}
int sfa_tmem = tmem_addr_sfa + 4 * sub_k_iter;
int sfb_tmem;
if constexpr (TD_MMA_N == 256) {
sfb_tmem = tmem_addr_sfb + 8 * sub_k_iter;
} else {
sfb_tmem = tmem_addr_sfb + 4 * sub_k_iter + (n_off%128)/32;
}
tcgen05_mma_nvfp4<1>(tmem_addr_result, a_desc, b_desc, make_instr_desc<TD_MMA_M, TD_MMA_N>(), sfa_tmem, sfb_tmem, k_off + sub_k_iter); // Inputting k_off like this will set enable-input-d so only on the first mma we 0 out the result space in TMEM
}
// signal MMA done
tcgen05_commit_multicast<CLUSTER_MASK>(mbar_addr_mma_stage);
glob_k_off += TD_SMEM_K;
}
// Signal epilogue to start
tcgen05_commit(mbar_addr_epi + 8*tmem_buf);
}
// All warps aside from the two for TMA/MMA are dedicated to the epilogue
else if (warp_id < NUM_WARPS - 3) {
const int M = M_vals[tile_stage];
const int tmem_addr_result = tmem_result_ptrs[tmem_buf];
mbar_wait(mbar_addr_epi + tmem_buf*8, epi_phase[tmem_buf]);
epi_phase[tmem_buf] ^= 1;
asm volatile("tcgen05.fence::after_thread_sync;");
constexpr int res_per_thrd = 4 * (OUT_N_CHUNK / 8);
float results[res_per_thrd];
constexpr int rows_per_warp = TD_MMA_M / (NUM_WARPS - 3);
for (int chunk = 0; chunk < TD_MMA_N / OUT_N_CHUNK; chunk++) {
int buf = chunk & 0x1; // faster chunk % 2
// Wait for previous TMA from this buffer to finish reading FIRST
if (chunk >= 2) {
if (warp_id == 0 && elect_one_sync()) {
asm volatile("cp.async.bulk.wait_group.read 1;");
}
// All threads must wait for the elected thread's wait to complete
asm volatile("bar.sync 2, %0;" :: "r"(WARP_SIZE * (NUM_WARPS - 3)) : "memory");
}
// Load chunk from TMEM -> REGS, cvt to half, store to SMEM
for (int sub_m = 0; sub_m < rows_per_warp / 16; sub_m++) {
if (m_off + warp_id * rows_per_warp + sub_m * 16 > M) {
break;
}
tcgen05_ld<16, 256, OUT_N_CHUNK / 8>(results, tmem_addr_result + (((warp_id * rows_per_warp) + sub_m * 16) << 16) + chunk * OUT_N_CHUNK);
asm volatile("tcgen05.wait::ld.sync.aligned;");
// Post process and store from Regs to SMEM (Regs -> SMEM)
// Transfer result from SMEM -> GMEM (8 comes from 256/32 -> 256b per ld block from above)
for (int i = 0; i < OUT_N_CHUNK / 8; i++) {
const int m_offset = warp_id * rows_per_warp + sub_m * 16 + lane_id / 4;
const int n_offset = i * 8 + (lane_id % 4) * 2;
if (m_offset + m_off < M) {
reinterpret_cast<half2 *>(c_smem + (C_CHUNK_SMEM_TILESZ * buf) + m_offset*OUT_N_CHUNK + n_offset)[0] = __float22half2_rn({results[i * 4], results[i * 4 + 1]});
}
if (m_offset + m_off + 8 < M) {
reinterpret_cast<half2 *>(c_smem + (C_CHUNK_SMEM_TILESZ * buf) + (m_offset + 8)*OUT_N_CHUNK + n_offset)[0] = __float22half2_rn({results[i * 4 + 2], results[i * 4 + 3]});
}
}
}
asm volatile("bar.sync 2, %0;" :: "r"(WARP_SIZE * (NUM_WARPS - 3)) : "memory");
// Only ONE thread issues TMA store and manages wait_group
if (warp_id == 0 && elect_one_sync()) {
asm volatile ("fence.proxy.async.shared::cta;" ::: "memory"); // ensure all SMEM writes complete and are visible to async proxy
tcgen05_2dtma_s2g_c(c_smem_ptr + (C_CHUNK_SMEM_TILESZ * 2 * buf), g_C_tmap + tile_stage, m_off, n_off + (chunk * OUT_N_CHUNK), CacheHintSm100::EVICT_NORMAL);
asm volatile("cp.async.bulk.commit_group;");
}
}
if (warp_id == 0 && elect_one_sync()) {
asm volatile("cp.async.bulk.wait_group.read 0;");
mbar_arrive(mbar_addr_epi_done + tmem_buf*8, 1);
}
tmem_buf ^= 1;
asm volatile("bar.sync 2, %0;" :: "r"(WARP_SIZE * (NUM_WARPS - 3)) : "memory");
}
if constexpr (!SINGLE_WAVE) {
if (warp_id == MMA_WARP) {
tmem_buf ^= 1;
if (tile_idx != blockIdx.x) {
mbar_wait(mbar_addr_epi_done + tmem_buf*8, epi_done_phase[tmem_buf]);
epi_done_phase[tmem_buf] ^= 1;
}
}
tile_stages++;
}
}
// Free memory
__syncthreads();
if (warp_id == 0) { tcgen05_dealloc_tmem<1>(tmem_addr_base, TD_MMA_N * 2 * 2); }
}
static CUtensorMap* d_tmaps;
static bool allocated = false;
static bool attr_set = false;
static PFN_cuTensorMapEncodeTiled_v12000 cuTensorMapEncodeTiled_fn;
static CUtensorMap tmap_a_temp, tmap_b_temp, tmap_c_temp;
static int cache_N = 0, cache_K = 0;
void nvfp4_group_gemm(const std::vector<std::tuple<torch::Tensor, torch::Tensor, torch::Tensor>>& abc_tensors, const std::vector<std::tuple<torch::Tensor, torch::Tensor>>& sf_tensors, const std::vector<std::tuple<int, int, int, int>>& prob_sizes, int N, int K, int G) {
// Constants
constexpr int M_TILE_SIZE = 128;
constexpr int K_MMA_SIZE = 64;
constexpr int NUM_WARPS = 7;
// Configurables
constexpr int N_TILE_SIZE = 128;
constexpr int K_TILE_SIZE = 256;
constexpr int PIPE_STAGES = 5;
constexpr int OUT_N_CHUNK = 32;
constexpr int TILE_DESC_PIPE_STAGES = 5;
constexpr int CLUSTER_SIZE = 2;
constexpr CUtensorMapSwizzle SWIZZLE_TYPE = CUtensorMapSwizzle::CU_TENSOR_MAP_SWIZZLE_128B;
// Build GroupDescs on the stack — passed directly as a kernel argument (constant memory)
GroupDescs gd = {};
bool nk_var = false;
int total_tiles = 0;
for (int i = 0; i < G; i++) {
gd.groups[i].A_addr = std::get<0>(abc_tensors[i]).data_ptr();
gd.groups[i].B_addr = std::get<1>(abc_tensors[i]).data_ptr();
gd.groups[i].C_addr = reinterpret_cast<__half*>(std::get<2>(abc_tensors[i]).data_ptr());
gd.groups[i].sfa_addr = reinterpret_cast<uint8_t*>(std::get<0>(sf_tensors[i]).data_ptr());
gd.groups[i].sfb_addr = reinterpret_cast<uint8_t*>(std::get<1>(sf_tensors[i]).data_ptr());
int M_size = std::get<0>(prob_sizes[i]);
int N_size = std::get<1>(prob_sizes[i]);
int K_size = std::get<2>(prob_sizes[i]);
gd.groups[i].M = M_size;
gd.groups[i].N = N_size;
gd.groups[i].K = K_size;
gd.groups[i].block_start = total_tiles;
total_tiles += CEIL_DIV(M_size, M_TILE_SIZE) * CEIL_DIV(N_size, N_TILE_SIZE);
if (N_size != N || K_size != K) {
nk_var = true;
}
}
bool SINGLE_WAVE = false;
int NUM_CTAS = 148;
if (total_tiles < 148) {
NUM_CTAS = total_tiles;
SINGLE_WAVE = true;
}
if (!allocated) {
cudaMalloc(&d_tmaps, 3 * 148 * sizeof(CUtensorMap) * TILE_DESC_PIPE_STAGES);
allocated = true;
}
// Query driver for tmap templates. Each CTA needs to fill in tmap_a_temp with the correct pointer / M value, and each tmap_b_temp with the correct pointer
if (N != cache_N || K != cache_K) {
cuTensorMapEncodeTiled_fn = get_cuTensorMapEncodeTiled();
tma_3d_map_ab<M_TILE_SIZE, K_TILE_SIZE, SWIZZLE_TYPE>::init(cuTensorMapEncodeTiled_fn, &tmap_a_temp, nullptr, M_TILE_SIZE, K);
tma_3d_map_ab<N_TILE_SIZE, K_TILE_SIZE, SWIZZLE_TYPE>::init(cuTensorMapEncodeTiled_fn, &tmap_b_temp, nullptr, N, K);
tma_2d_map_c_init<M_TILE_SIZE, OUT_N_CHUNK>(cuTensorMapEncodeTiled_fn, &tmap_c_temp, nullptr, M_TILE_SIZE, N);
cache_N = N;
cache_K = K;
}
// === KERNEL LAUNCH ===
auto kernel_multiwave_nkconst = nvfp4_group_gemm_kernel<M_TILE_SIZE, N_TILE_SIZE, M_TILE_SIZE, N_TILE_SIZE, K_TILE_SIZE, M_TILE_SIZE, N_TILE_SIZE, K_MMA_SIZE, SWIZZLE_TYPE, PIPE_STAGES, NUM_WARPS, OUT_N_CHUNK, TILE_DESC_PIPE_STAGES, CLUSTER_SIZE, false, false>;
auto kernel_singlewave_nkconst = nvfp4_group_gemm_kernel<M_TILE_SIZE, N_TILE_SIZE, M_TILE_SIZE, N_TILE_SIZE, K_TILE_SIZE, M_TILE_SIZE, N_TILE_SIZE, K_MMA_SIZE, SWIZZLE_TYPE, PIPE_STAGES, NUM_WARPS, OUT_N_CHUNK, TILE_DESC_PIPE_STAGES, CLUSTER_SIZE, true, false>;
auto kernel_multiwave_nkvar = nvfp4_group_gemm_kernel<M_TILE_SIZE, N_TILE_SIZE, M_TILE_SIZE, N_TILE_SIZE, K_TILE_SIZE, M_TILE_SIZE, N_TILE_SIZE, K_MMA_SIZE, SWIZZLE_TYPE, PIPE_STAGES, NUM_WARPS, OUT_N_CHUNK, TILE_DESC_PIPE_STAGES, 1, false, true>;
auto kernel_singlewave_nkvar = nvfp4_group_gemm_kernel<M_TILE_SIZE, N_TILE_SIZE, M_TILE_SIZE, N_TILE_SIZE, K_TILE_SIZE, M_TILE_SIZE, N_TILE_SIZE, K_MMA_SIZE, SWIZZLE_TYPE, PIPE_STAGES, NUM_WARPS, OUT_N_CHUNK, TILE_DESC_PIPE_STAGES, 1, true, true>;
if (!attr_set) {
cudaFuncSetAttribute(
kernel_multiwave_nkconst,
cudaFuncAttributePreferredSharedMemoryCarveout,
cudaSharedmemCarveoutMaxShared // Maximum shared memory
);
cudaFuncSetAttribute(
kernel_singlewave_nkconst,
cudaFuncAttributePreferredSharedMemoryCarveout,
cudaSharedmemCarveoutMaxShared // Maximum shared memory
);
cudaFuncSetAttribute(
kernel_multiwave_nkvar,
cudaFuncAttributePreferredSharedMemoryCarveout,
cudaSharedmemCarveoutMaxShared // Maximum shared memory
);
cudaFuncSetAttribute(
kernel_singlewave_nkvar,
cudaFuncAttributePreferredSharedMemoryCarveout,
cudaSharedmemCarveoutMaxShared // Maximum shared memory
);
attr_set = true;
}
constexpr int threads = WARP_SIZE * NUM_WARPS;
if (SINGLE_WAVE) {
if (nk_var) {
kernel_singlewave_nkvar<<<NUM_CTAS, threads>>>(gd, tmap_a_temp, tmap_b_temp, tmap_c_temp, total_tiles, d_tmaps, N, K, G);
} else {
kernel_singlewave_nkconst<<<NUM_CTAS, threads>>>(gd, tmap_a_temp, tmap_b_temp, tmap_c_temp, total_tiles, d_tmaps, N, K, G);
}
} else {
if (nk_var) {
kernel_multiwave_nkvar<<<NUM_CTAS, threads>>>(gd, tmap_a_temp, tmap_b_temp, tmap_c_temp, total_tiles, d_tmaps, N, K, G);
} else {
kernel_multiwave_nkconst<<<NUM_CTAS, threads>>>(gd, tmap_a_temp, tmap_b_temp, tmap_c_temp, total_tiles, d_tmaps, N, K, G);
}
}
}
"""
nvfp4_group_gemm_cpp_source = """
#include <torch/extension.h>
void nvfp4_group_gemm(const std::vector<std::tuple<torch::Tensor, torch::Tensor, torch::Tensor>>& abc_tensors, const std::vector<std::tuple<torch::Tensor, torch::Tensor>>& sf_tensors, const std::vector<std::tuple<int, int, int, int>>& prob_sizes, int N, int K, int G);
"""
nvfp4_group_gemm_module = load_inline(
name='nvfp4_group_gemm',
cpp_sources=nvfp4_group_gemm_cpp_source,
cuda_sources=nvfp4_group_gemm_cuda_source,
functions=['nvfp4_group_gemm'],
verbose=True,
extra_cuda_cflags=[
'-I/usr/local/lib/python3.12/site-packages/cutlass_library/source/include',
'-I/usr/local/lib/python3.12/site-packages/cutlass_library/source/tools/util/include',
'-O3',
'-gencode=arch=compute_100a,code=sm_100a',
'-Xptxas', '--allow-expensive-optimizations=true',
'--use_fast_math',
'--relocatable-device-code=false',
# '-lineinfo',
],
)
def custom_kernel(data: input_t) -> output_t:
import time
abc_tensors, _, sfasfb_tensors_reordered, problem_sizes = data
"""
abc_tensors: [(a_ref, b_ref, c_ref), (a_ref, b_ref, c_ref), ...]
sfasfb_tensors_reordered: [(sfa, sfb), (sfa, sfb), ...]
problem_sizes (l is always 1): [(m, n, k, l), (m, n, k, l), ...]
"""
G = len(abc_tensors)
N = problem_sizes[0][1]
K = problem_sizes[0][2]
nvfp4_group_gemm_module.nvfp4_group_gemm(abc_tensors, sfasfb_tensors_reordered, problem_sizes, N, K, G)
return [c for (a, b, c) in abc_tensors]
scrolls · 1304 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