submission 377876
nrehiew · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 1023 lines, June 9 Researcher Reciprocity License v1.0.
submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-modal-nvfp4-dual-gemm-377876?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:9114eb0b12174c295cbbdbdd2c949d84f133271163db5040c7ed843ccc4f9efb
license declaredunknown
license concludedunknown
authorsnrehiew
imported2026-08-15
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
async-copy
__device__ __forceinline__ void cp_async_bulk_gmem2smem(void* dst, const void* src, uint32_t size, uint64_t* mbar) {cluster
__cluster_dims__(2, 1, 1)fused-epilogue
TAG_EPILOGUE = 4, // Time for writeback to global memorymbarrier
asm volatile("tcgen05.commit.cta_group::1.mbarrier::arrive::one.b64 [%0];"num-warps = 8
constexpr int NUM_WARPS = 8;shared-memory
__device__ __forceinline__ void cp_async_bulk_tensor_gmem2smem_multicast(void* dst, const CUtensorMap* tmap, const int32_t* coords, uint64_t* mbar) {tcgen05
"tcgen05.alloc.cta_group::1.sync.aligned.b32 [%0], %1;\n"tile-n = 64
constexpr int STORE_BLOCK_N = 64; // 128 bytes / 2 bytes per bf16 for 128B swizzletma
asm volatile("cp.async.bulk.shared::cta.global.mbarrier::complete_tx::bytes.L2::cache_hint [%0], [%1], %2, [%3], %4;"vector-width = half2
half2 packed[4];Kernel source
submission.py1023 lines
# popcorn-cli submit submission.py --no-tui --leaderboard nvfp4_dual_gemm --gpu NVIDIA --mode test
from task import input_t, output_t
from torch.utils.cpp_extension import load_inline
cpp_source = """
#include <torch/extension.h>
torch::Tensor cuda_entry(torch::Tensor A, torch::Tensor B1, torch::Tensor B2, torch::Tensor SFA_permuted, torch::Tensor SFB1_permuted, torch::Tensor SFB2_permuted, torch::Tensor C);
"""
CUDA_SRC = r"""
#include <cuda.h>
#include <torch/torch.h>
#include <cuda_bf16.h>
#include <cooperative_groups.h>
#include <cuda/barrier>
#include <cuda/ptx>
#include <cudaTypedefs.h>
#include <cuda.h>
#include <cstdint>
namespace ptx = cuda::ptx;
#define PROFILE_TAGS_DEFINED
enum GemmProfileTags {
TAG_TMA_ISSUE = 0, // Time to issue TMA commands
TAG_TMA_WAIT = 1, // Time waiting for TMA
TAG_MMA_ISSUE = 2, // Time to issue MMA commands
TAG_MMA_WAIT = 3, // Time waiting for MMA
TAG_EPILOGUE = 4, // Time for writeback to global memory
TAG_SCALE_LOADING = 5, // Time for scale loading
TAG_SETUP = 6, // Time for setup
};
#ifndef PROFILE_MAX_EVENTS
#define PROFILE_MAX_EVENTS 65536
#endif
#define PROFILE_BUFFER_SIZE (1 + PROFILE_MAX_EVENTS * 4)
// Buffer layout: profile[0] = count, then for event i:
// profile[1+i*4+0] = start_ns, [1+i*4+1] = duration_ns, [1+i*4+2] = tag, [1+i*4+3] = tid
struct IntraKernelProfiler {
int64_t* profile;
int current_event_id;
int tid;
__device__ __forceinline__ void init(int64_t* profile_buffer, int thread_id) {
profile = profile_buffer;
current_event_id = -1;
tid = thread_id;
}
__device__ __forceinline__ int64_t read_globaltimer() {
int64_t time;
asm volatile("mov.u64 %0, %%globaltimer;" : "=l"(time));
return time;
}
__device__ __forceinline__ void start(bool condition = true) {
if (!condition) {
current_event_id = -1;
return;
}
int i = atomicAdd(reinterpret_cast<int*>(profile), 1);
if (i >= PROFILE_MAX_EVENTS) {
current_event_id = -1;
return;
}
current_event_id = i;
profile[1 + i * 4] = read_globaltimer();
}
__device__ __forceinline__ void stop(int tag, bool condition = true) {
if (!condition || current_event_id < 0 || current_event_id >= PROFILE_MAX_EVENTS) return;
int64_t end_time = read_globaltimer();
profile[1 + current_event_id * 4 + 1] = end_time - profile[1 + current_event_id * 4];
profile[1 + current_event_id * 4 + 2] = tag;
profile[1 + current_event_id * 4 + 3] = tid;
}
};
// Cache hint constants for L2 cache policy
// 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;
constexpr int SF_VEC_SIZE = 16;
constexpr int ELEMENTS_PER_BYTE = 2;
constexpr int MMA_K = 64; // nvf4 processes 64 elements per MMA
constexpr int MMA_K_SCALES = MMA_K / SF_VEC_SIZE; // 4
constexpr int SF_NUM_COLS_PER_ITER = 4;
constexpr int MMA_K_IN_BYTES = MMA_K / ELEMENTS_PER_BYTE; // 32 bytes
constexpr int K_LOAD_SIZE_IN_BYTES = 128;
constexpr int WARP_SIZE = 32;
constexpr int NUM_WARPS = 8;
constexpr int NUM_THREADS = NUM_WARPS * WARP_SIZE;
constexpr int SF_TILE_SIZE_BYTES = 512;
constexpr int STORE_BLOCK_N = 64; // 128 bytes / 2 bytes per bf16 for 128B swizzle
// https://github.com/NVIDIA/cutlass/blob/v4.2.1/include/cute/arch/cluster_sm90.hpp#L180
__device__ inline
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;
}
inline unsigned int cdiv(unsigned int a, unsigned int b) {
return (a + b - 1) / b;
}
__device__ __forceinline__
uint32_t cvta_shared(const void *ptr) {
return static_cast<uint32_t>(__cvta_generic_to_shared(ptr));
}
__forceinline__ __device__ uint32_t get_tmem_addr(uint32_t base_addr, int row_offset, int col_offset) {
uint32_t offset = (row_offset << 16) | (col_offset & 0xFFFF);
return base_addr + offset;
}
// PTX
__device__ inline void allocate_tensor_memory(uint32_t* ptr, int n_cols) {
asm volatile(
"tcgen05.alloc.cta_group::1.sync.aligned.b32 [%0], %1;\n"
:: "l"(ptr), "r"(n_cols)
);
}
__device__ inline void deallocate_tensor_memory(uint32_t tmem_addr, int n_cols) {
asm volatile(
"tcgen05.dealloc.cta_group::1.sync.aligned.b32 %0, %1;\n"
:: "r"(tmem_addr), "r"(n_cols)
);
}
__device__ __forceinline__ void tcgen05_commit_group(uint64_t *bar) {
uint32_t mbar_ptr = static_cast<uint32_t>(__cvta_generic_to_shared(bar));
asm volatile("tcgen05.commit.cta_group::1.mbarrier::arrive::one.b64 [%0];"
:: "r"(mbar_ptr) // 32-bit shared address
);
}
__device__ static __forceinline__ void init_barrier(uint64_t* bar, int thread_count) {
uint32_t bar_ptr = static_cast<uint32_t>(__cvta_generic_to_shared(bar));
asm volatile (
"mbarrier.init.shared::cta.b64 [%0], %1;\n"
:: "r"(bar_ptr), "r"(thread_count)
);
}
__device__ static __forceinline__ void expect_bytes(uint64_t* bar, uint32_t bytes) {
uint32_t bar_ptr = static_cast<uint32_t>(__cvta_generic_to_shared(bar));
asm volatile(
"mbarrier.arrive.expect_tx.shared::cta.b64 _, [%0], %1;\n"
:: "r"(bar_ptr), "r"(bytes)
);
}
__device__ static __forceinline__ void wait(uint64_t& bar, int phase_bit) {
uint32_t mbar_ptr = static_cast<uint32_t>(__cvta_generic_to_shared(&bar));
asm volatile (
"{\n"
".reg .pred P1;\n"
"LAB_WAIT:\n"
"mbarrier.try_wait.parity.acquire.cta.shared::cta.b64 P1, [%0], %1;\n"
"@P1 bra.uni DONE;\n"
"bra.uni LAB_WAIT;\n"
"DONE:\n"
"}\n"
:: "r"(mbar_ptr), "r"(phase_bit)
);
}
template<uint64_t CACHE_POLICY>
__device__ __forceinline__ void cp_async_bulk_gmem2smem(void* dst, const void* src, uint32_t size, uint64_t* mbar) {
uint32_t dst_addr = static_cast<uint32_t>(__cvta_generic_to_shared(dst));
uint32_t mbar_addr = static_cast<uint32_t>(__cvta_generic_to_shared(mbar));
asm volatile("cp.async.bulk.shared::cta.global.mbarrier::complete_tx::bytes.L2::cache_hint [%0], [%1], %2, [%3], %4;"
:: "r"(dst_addr), "l"(src), "r"(size), "r"(mbar_addr), "l"(CACHE_POLICY));
}
template<uint64_t CACHE_POLICY>
__device__ __forceinline__ void cp_async_bulk_tensor_gmem2smem(void* dst, const CUtensorMap* tmap, const int32_t* coords, uint64_t* mbar) {
uint32_t dst_addr = static_cast<uint32_t>(__cvta_generic_to_shared(dst));
uint32_t mbar_addr = static_cast<uint32_t>(__cvta_generic_to_shared(mbar));
asm volatile("cp.async.bulk.tensor.3d.shared::cta.global.mbarrier::complete_tx::bytes.cta_group::1.L2::cache_hint "
"[%0], [%1, {%2, %3, %4}], [%5], %6;"
:: "r"(dst_addr), "l"(tmap), "r"(coords[0]), "r"(coords[1]), "r"(coords[2]), "r"(mbar_addr), "l"(CACHE_POLICY)
: "memory");
}
template<uint64_t CACHE_POLICY>
__device__ __forceinline__ void cp_async_bulk_tensor_gmem2smem_multicast(void* dst, const CUtensorMap* tmap, const int32_t* coords, uint64_t* mbar) {
uint32_t dst_addr = static_cast<uint32_t>(__cvta_generic_to_shared(dst));
uint32_t mbar_addr = static_cast<uint32_t>(__cvta_generic_to_shared(mbar));
asm volatile("cp.async.bulk.tensor.3d.shared::cluster.global.mbarrier::complete_tx::bytes.cta_group::1.L2::cache_hint "
"[%0], [%1, {%2, %3, %4}], [%5], %6;"
:: "r"(dst_addr), "l"(tmap), "r"(coords[0]), "r"(coords[1]), "r"(coords[2]), "r"(mbar_addr), "l"(CACHE_POLICY)
: "memory");
}
__device__ __forceinline__ void st_shared(const void* ptr, uint32_t x, uint32_t y, uint32_t z, uint32_t w) {
asm volatile("st.shared.v4.u32 [%0], {%1, %2, %3, %4};" :: "l"(__cvta_generic_to_shared(ptr)), "r"(x), "r"(y), "r"(z), "r"(w));
}
__device__ inline
uint64_t matrix_descriptor_encode(const uint64_t index) {
return (index & 0x3FFFF) >> 4;
}
__device__
uint64_t make_matrix_descriptor_sf(void* x) {
constexpr int sbo = 16 * 8;
constexpr int lbo = 16;
// Blackwell reference 9.7.16.4.1
// Both A and B are contiguous along K (K Major)
uint64_t descriptor = 0;
descriptor |= matrix_descriptor_encode(cvta_shared(x)); // Matrix start address
// leading dimension byte offset (row size in bytes for K-major layout)
// Only used when swizzle_mode = 0 (no swizzle)
descriptor |= matrix_descriptor_encode(lbo) << 16;
// stride dimension byte offset (offset from the first 8 columns to the next 8 columns)
// For FP4 packed: stride = 8 rows × K_BYTES bytes/row
descriptor |= matrix_descriptor_encode(sbo) << 32;
descriptor |= 0b001ULL << 46; // Fixed constant value of 0b001
return descriptor;
}
__device__
uint64_t make_matrix_descriptor_matmul(void* x) {
// Blackwell reference 9.7.16.4.1
// Both A and B are contiguous along K (K Major)
uint64_t descriptor = 0;
constexpr int sbo = K_LOAD_SIZE_IN_BYTES * 8;
descriptor |= matrix_descriptor_encode(cvta_shared(x)); // Matrix start address
// leading dimension byte offset (row size in bytes for K-major layout)
// stride dimension byte offset (offset from the first 8 columns to the next 8 columns)
// For FP4 packed: stride = 8 rows × K_BYTES bytes/row
descriptor |= matrix_descriptor_encode(sbo) << 32;
descriptor |= 0b001ULL << 46; // Fixed constant value of 0b001
descriptor |= (2llu) << 61; // 2. 128-Byte swizzling
return descriptor;
}
__device__ __forceinline__ void tcgen05_ld_128(float* results, uint32_t base_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"(results[0]), "=f"(results[1]), "=f"(results[2]), "=f"(results[3]), "=f"(results[4]), "=f"(results[5]), "=f"(results[6]), "=f"(results[7]),
"=f"(results[8]), "=f"(results[9]), "=f"(results[10]), "=f"(results[11]), "=f"(results[12]), "=f"(results[13]), "=f"(results[14]), "=f"(results[15]),
"=f"(results[16]), "=f"(results[17]), "=f"(results[18]), "=f"(results[19]), "=f"(results[20]), "=f"(results[21]), "=f"(results[22]), "=f"(results[23]),
"=f"(results[24]), "=f"(results[25]), "=f"(results[26]), "=f"(results[27]), "=f"(results[28]), "=f"(results[29]), "=f"(results[30]), "=f"(results[31]),
"=f"(results[32]), "=f"(results[33]), "=f"(results[34]), "=f"(results[35]), "=f"(results[36]), "=f"(results[37]), "=f"(results[38]), "=f"(results[39]),
"=f"(results[40]), "=f"(results[41]), "=f"(results[42]), "=f"(results[43]), "=f"(results[44]), "=f"(results[45]), "=f"(results[46]), "=f"(results[47]),
"=f"(results[48]), "=f"(results[49]), "=f"(results[50]), "=f"(results[51]), "=f"(results[52]), "=f"(results[53]), "=f"(results[54]), "=f"(results[55]),
"=f"(results[56]), "=f"(results[57]), "=f"(results[58]), "=f"(results[59]), "=f"(results[60]), "=f"(results[61]), "=f"(results[62]), "=f"(results[63])
: "r"(base_addr)
);
}
__device__ __forceinline__ void tcgen05_ld_64(float* results, uint32_t base_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"(results[0]), "=f"(results[1]), "=f"(results[2]), "=f"(results[3]), "=f"(results[4]), "=f"(results[5]), "=f"(results[6]), "=f"(results[7]),
"=f"(results[8]), "=f"(results[9]), "=f"(results[10]), "=f"(results[11]), "=f"(results[12]), "=f"(results[13]), "=f"(results[14]), "=f"(results[15]),
"=f"(results[16]), "=f"(results[17]), "=f"(results[18]), "=f"(results[19]), "=f"(results[20]), "=f"(results[21]), "=f"(results[22]), "=f"(results[23]),
"=f"(results[24]), "=f"(results[25]), "=f"(results[26]), "=f"(results[27]), "=f"(results[28]), "=f"(results[29]), "=f"(results[30]), "=f"(results[31])
: "r"(base_addr)
);
}
template <const int BLOCK_N>
__device__ __forceinline__ void tcgen05_ld(float* results, uint32_t base_addr){
// results is always BLOCK_N/2 registers
if constexpr (BLOCK_N == 128) {
tcgen05_ld_128(results, base_addr);
} else if constexpr (BLOCK_N == 64) {
tcgen05_ld_64(results, base_addr);
} else {
static_assert(BLOCK_N != 128 && BLOCK_N != 64, "Invalid block size");
}
}
__device__ __forceinline__ void tcgen05_ld_8(float* results, uint32_t base_addr) {
asm volatile(
"tcgen05.ld.sync.aligned.32x32b.x8.b32 {%0,%1,%2,%3,%4,%5,%6,%7}, [%8];"
: "=f"(results[0]), "=f"(results[1]), "=f"(results[2]), "=f"(results[3]),
"=f"(results[4]), "=f"(results[5]), "=f"(results[6]), "=f"(results[7])
: "r"(base_addr)
);
}
__device__ __forceinline__ void tcgen05_cp(uint32_t tmem_addr, uint64_t sdesc) {
asm volatile(
"tcgen05.cp.cta_group::1.32x128b.warpx4 [%0], %1;" // shape: [no. lanes in TMEM, size in bits across columns]
:: "r"(tmem_addr), "l"(sdesc)
);
}
// nvf4 MMA with block scaling - matrices in shared memory, scale factors in tensor memory
template <bool init>
__device__ __forceinline__ void tcgen05_mma_nvf4(
uint32_t d_tmem_addr,
uint32_t i_desc,
void* sA,
void* sB,
uint32_t scale_a_tmem,
uint32_t scale_b_tmem
) {
uint64_t a_desc = make_matrix_descriptor_matmul(sA);
uint64_t b_desc = make_matrix_descriptor_matmul(sB);
if (init) {
asm volatile(
"{\n"
".reg .pred p;\n"
"setp.eq.u32 p, 1, 0;\n" // p = False (initialize)
"tcgen05.mma.cta_group::1.kind::mxf4nvf4.block_scale.block16 "
"[%0], %1, %2, %3, [%4], [%5], p;\n"
"}\n"
:: "r"(d_tmem_addr), "l"(a_desc), "l"(b_desc), "r"(i_desc),
"r"(scale_a_tmem), "r"(scale_b_tmem)
);
} else {
asm volatile(
"{\n"
".reg .pred p;\n"
"setp.eq.u32 p, 1, 1;\n" // p = True (accumulate)
"tcgen05.mma.cta_group::1.kind::mxf4nvf4.block_scale.block16 "
"[%0], %1, %2, %3, [%4], [%5], p;\n"
"}\n"
:: "r"(d_tmem_addr), "l"(a_desc), "l"(b_desc), "r"(i_desc),
"r"(scale_a_tmem), "r"(scale_b_tmem)
);
}
}
// nvf4 MMA with collector buffer fill - fills collector with matrix A
template <bool init>
__device__ __forceinline__ void tcgen05_mma_nvf4_fill(
uint32_t d_tmem_addr,
uint32_t i_desc,
void* sA,
void* sB,
uint32_t scale_a_tmem,
uint32_t scale_b_tmem
) {
uint64_t a_desc = make_matrix_descriptor_matmul(sA);
uint64_t b_desc = make_matrix_descriptor_matmul(sB);
if (init) {
asm volatile(
"{\n"
".reg .pred p;\n"
"setp.eq.u32 p, 1, 0;\n" // p = False (initialize)
"tcgen05.mma.cta_group::1.kind::mxf4nvf4.block_scale.block16.collector::a::fill "
"[%0], %1, %2, %3, [%4], [%5], p;\n"
"}\n"
:: "r"(d_tmem_addr), "l"(a_desc), "l"(b_desc), "r"(i_desc),
"r"(scale_a_tmem), "r"(scale_b_tmem)
);
} else {
asm volatile(
"{\n"
".reg .pred p;\n"
"setp.eq.u32 p, 1, 1;\n" // p = True (accumulate)
"tcgen05.mma.cta_group::1.kind::mxf4nvf4.block_scale.block16.collector::a::fill "
"[%0], %1, %2, %3, [%4], [%5], p;\n"
"}\n"
:: "r"(d_tmem_addr), "l"(a_desc), "l"(b_desc), "r"(i_desc),
"r"(scale_a_tmem), "r"(scale_b_tmem)
);
}
}
// nvf4 MMA with collector buffer lastuse - uses cached A from collector and discards
template <bool init>
__device__ __forceinline__ void tcgen05_mma_nvf4_lastuse(
uint32_t d_tmem_addr,
uint32_t i_desc,
void* sA,
void* sB,
uint32_t scale_a_tmem,
uint32_t scale_b_tmem
) {
uint64_t a_desc = make_matrix_descriptor_matmul(sA);
uint64_t b_desc = make_matrix_descriptor_matmul(sB);
if (init) {
asm volatile(
"{\n"
".reg .pred p;\n"
"setp.eq.u32 p, 1, 0;\n" // p = False (initialize)
"tcgen05.mma.cta_group::1.kind::mxf4nvf4.block_scale.block16.collector::a::lastuse "
"[%0], %1, %2, %3, [%4], [%5], p;\n"
"}\n"
:: "r"(d_tmem_addr), "l"(a_desc), "l"(b_desc), "r"(i_desc),
"r"(scale_a_tmem), "r"(scale_b_tmem)
);
} else {
asm volatile(
"{\n"
".reg .pred p;\n"
"setp.eq.u32 p, 1, 1;\n" // p = True (accumulate)
"tcgen05.mma.cta_group::1.kind::mxf4nvf4.block_scale.block16.collector::a::lastuse "
"[%0], %1, %2, %3, [%4], [%5], p;\n"
"}\n"
:: "r"(d_tmem_addr), "l"(a_desc), "l"(b_desc), "r"(i_desc),
"r"(scale_a_tmem), "r"(scale_b_tmem)
);
}
}
__device__ static __forceinline__ void named_barrier_sync(int barrier_id, int thread_count) {
asm volatile("barrier.sync %0, %1;" :: "r"(barrier_id), "r"(thread_count));
}
__device__ __forceinline__ float silu(float x) {
// return x / (1.0f + expf(-x));
return x * (1.0f / (1.0f + __expf(-x)));
}
constexpr size_t align_up(size_t x, size_t a) {
return (x + a - 1) & ~(a - 1);
}
template<const int BLOCK_M, const int BLOCK_N, const int BLOCK_K, const int NUM_STAGES, const bool USE_TMA_EPILOGUE, const bool DO_PROFILE>
__launch_bounds__(NUM_THREADS)
__cluster_dims__(2, 1, 1)
__global__ void cuda_kernel(
const int M,
const int N,
const int K,
const __grid_constant__ CUtensorMap a_tensor_map,
const __grid_constant__ CUtensorMap b1_tensor_map,
const __grid_constant__ CUtensorMap b2_tensor_map,
const __grid_constant__ CUtensorMap c_tensor_map,
char* sfa_ptr,
char* sfb1_ptr,
char* sfb2_ptr,
half* __restrict__ C,
int64_t* profile_buf
){
constexpr int BLOCK_K_IN_BYTES = BLOCK_K / ELEMENTS_PER_BYTE;
constexpr int NUM_MMA_ITERS = BLOCK_K / MMA_K;
constexpr int SFA_COLS = SF_NUM_COLS_PER_ITER * NUM_MMA_ITERS;
constexpr int SFB_COLS = SF_NUM_COLS_PER_ITER * NUM_MMA_ITERS;
constexpr int A_bytes = BLOCK_M * BLOCK_K_IN_BYTES;
constexpr int B_bytes = BLOCK_N * BLOCK_K_IN_BYTES;
constexpr int SF_bytes = SF_TILE_SIZE_BYTES * NUM_MMA_ITERS;
constexpr int matrix_bytes = A_bytes + B_bytes * 2; // A, B1, B2
constexpr int sf_tma_bytes = SF_bytes * 3; // SFA, SFB1, SFB2
constexpr int stage_size = matrix_bytes + sf_tma_bytes;
constexpr int shared_size = stage_size * NUM_STAGES;
constexpr int SFA_PER_TILE_SIZE = BLOCK_M * MMA_K_SCALES; // for one mma iter
constexpr int SFB_PER_TILE_SIZE = SFA_PER_TILE_SIZE; // for one mma iter
constexpr uint32_t idesc_nvf4 =
(1 << 7) | (1 << 10) | ((BLOCK_N >> 3) << 17) | ((BLOCK_M >> 7) << 27);
const int warp_id = threadIdx.x / WARP_SIZE;
const int block_idx = blockIdx.x;
uint32_t cluster_ctaid;
asm volatile("mov.b32 %0, %%cluster_ctarank;" : "=r"(cluster_ctaid));
// const int block_row = block_idx / (N / BLOCK_N);
// const int block_col = block_idx % (N / BLOCK_N);
const int cluster_idx = blockIdx.x / 2;
const int blocks_per_row = N / BLOCK_N;
const int cluster_row = (cluster_idx / blocks_per_row) * 2;
const int cluster_col = cluster_idx % blocks_per_row;
const int block_row = cluster_row + cluster_ctaid;
const int block_col = cluster_col;
IntraKernelProfiler profiler;
bool should_profile = false;
if constexpr (DO_PROFILE) {
profiler.init(profile_buf, block_idx * 10 + warp_id);
should_profile = (blockIdx.x < 20) && (threadIdx.x % WARP_SIZE == 0);
}
if constexpr (DO_PROFILE) profiler.start(should_profile);
C += block_row * BLOCK_M * N + block_col * BLOCK_N;
extern __shared__ __align__(128) uint8_t shared_mem[];
uint8_t* smem_ptr = shared_mem;
uint8_t* As = smem_ptr;
smem_ptr += align_up(A_bytes * NUM_STAGES, 128);
uint8_t* B1s = smem_ptr;
smem_ptr += align_up(B_bytes * NUM_STAGES, 128);
uint8_t* B2s = smem_ptr;
smem_ptr += align_up(B_bytes * NUM_STAGES, 128);
uint8_t* SFA = smem_ptr;
smem_ptr += align_up(SF_bytes * NUM_STAGES, 128);
uint8_t* SFB1 = smem_ptr;
smem_ptr += align_up(SF_bytes * NUM_STAGES, 128);
uint8_t* SFB2 = smem_ptr;
smem_ptr += align_up(SF_bytes * NUM_STAGES, 128);
half* Cs = reinterpret_cast<half*>(As);
__shared__ uint32_t tmem_addr_base_shared;
if (warp_id == 1) {
allocate_tensor_memory(&tmem_addr_base_shared, 512);
}
// Tensor memory layout: SFA * NUM_STAGES, SFB1 * NUM_STAGES, SFB2 * NUM_STAGES
constexpr int SF_TOTAL_COLS = SFA_COLS * NUM_STAGES + SFB_COLS * NUM_STAGES * 2;
__shared__ __align__(8) uint64_t is_empty_bars[NUM_STAGES];
__shared__ __align__(8) uint64_t mma_ready_bars[NUM_STAGES];
__shared__ __align__(8) uint64_t sf_tma_ready_bars[NUM_STAGES];
__shared__ __align__(8) uint64_t last_stage_bar;
if (threadIdx.x < NUM_STAGES) {
init_barrier(&is_empty_bars[threadIdx.x], 2);
init_barrier(&mma_ready_bars[threadIdx.x], 2);
init_barrier(&sf_tma_ready_bars[threadIdx.x], 1);
}
if (threadIdx.x == 0) {
init_barrier(&last_stage_bar, 2);
}
__syncthreads();
const uint32_t tmem_addr_base = tmem_addr_base_shared;
const uint32_t sfa_tmem_addr = get_tmem_addr(tmem_addr_base, 0, 0);
const uint32_t sfb1_tmem_addr = get_tmem_addr(sfa_tmem_addr, 0, SFA_COLS * NUM_STAGES);
const uint32_t sfb2_tmem_addr = get_tmem_addr(sfb1_tmem_addr, 0, SFB_COLS * NUM_STAGES);
const uint32_t d1_tmem_addr = get_tmem_addr(sfb2_tmem_addr, 0, SFB_COLS * NUM_STAGES);
const uint32_t d2_tmem_addr = get_tmem_addr(tmem_addr_base, 0, 512 - BLOCK_N);
int row_start = block_row * BLOCK_M;
int col_start = block_col * BLOCK_N;
int K_BYTES = K / 2;
int rest_k = K / SF_VEC_SIZE / 4;
int num_iters = K_BYTES / BLOCK_K_IN_BYTES;
sfa_ptr += row_start/128 * SF_TILE_SIZE_BYTES * rest_k;
sfb1_ptr += col_start/128 * SF_TILE_SIZE_BYTES * rest_k;
sfb2_ptr += col_start/128 * SF_TILE_SIZE_BYTES * rest_k;
if constexpr (DO_PROFILE) profiler.stop(TAG_SETUP, should_profile);
if (warp_id == 4 && elect_sync()) {
auto issue_tma = [&](int iter, int stage_id) {
if constexpr (DO_PROFILE) profiler.start(should_profile);
constexpr int NUM_K_CHUNKS = BLOCK_K_IN_BYTES / K_LOAD_SIZE_IN_BYTES;
const int k_chunk_start = iter * NUM_K_CHUNKS;
int32_t a_tensor_coords[3] = {0, row_start, k_chunk_start};
int32_t b12_tensor_coords[3] = {0, col_start, k_chunk_start};
cp_async_bulk_gmem2smem<EVICT_FIRST>(&SFA[stage_id * SF_bytes], sfa_ptr + iter * SF_bytes, SF_bytes, &sf_tma_ready_bars[stage_id]);
cp_async_bulk_gmem2smem<EVICT_FIRST>(&SFB1[stage_id * SF_bytes], sfb1_ptr + iter * SF_bytes, SF_bytes, &sf_tma_ready_bars[stage_id]);
cp_async_bulk_gmem2smem<EVICT_FIRST>(&SFB2[stage_id * SF_bytes], sfb2_ptr + iter * SF_bytes, SF_bytes, &sf_tma_ready_bars[stage_id]);
expect_bytes(&sf_tma_ready_bars[stage_id], sf_tma_bytes);
cp_async_bulk_tensor_gmem2smem<EVICT_LAST>(&As[stage_id * A_bytes], &a_tensor_map, a_tensor_coords, &mma_ready_bars[stage_id]);
cp_async_bulk_tensor_gmem2smem<EVICT_NORMAL>(&B1s[stage_id * B_bytes], &b1_tensor_map, b12_tensor_coords, &mma_ready_bars[stage_id]);
cp_async_bulk_tensor_gmem2smem<EVICT_NORMAL>(&B2s[stage_id * B_bytes], &b2_tensor_map, b12_tensor_coords, &mma_ready_bars[stage_id]);
expect_bytes(&mma_ready_bars[stage_id], matrix_bytes);
if constexpr (DO_PROFILE) profiler.stop(TAG_TMA_ISSUE, should_profile);
};
#pragma unroll
for (int iter = 0; iter < NUM_STAGES; iter++) {
issue_tma(iter, iter);
}
#pragma unroll
for (int iter = NUM_STAGES; iter < num_iters; iter++) {
const int stage_id = iter % NUM_STAGES;
const int mma_phase = (iter / NUM_STAGES - 1) % 2;
if constexpr (DO_PROFILE) profiler.start(should_profile);
wait(is_empty_bars[stage_id], mma_phase);
if constexpr (DO_PROFILE) profiler.stop(TAG_MMA_WAIT, should_profile);
issue_tma(iter, stage_id);
}
} else if (warp_id == 5 && elect_sync()) {
uint32_t sf_phase[NUM_STAGES] = {};
int current_stage_idx = 0;
for (int iter = 0; iter < num_iters; iter++) {
uint32_t sfa_stage_base = get_tmem_addr(sfa_tmem_addr, 0, current_stage_idx * SFA_COLS);
uint32_t sfb1_stage_base = get_tmem_addr(sfb1_tmem_addr, 0, current_stage_idx * SFB_COLS);
uint32_t sfb2_stage_base = get_tmem_addr(sfb2_tmem_addr, 0, current_stage_idx * SFB_COLS);
if constexpr (DO_PROFILE) profiler.start(should_profile);
wait(sf_tma_ready_bars[current_stage_idx], sf_phase[current_stage_idx]);
sf_phase[current_stage_idx] ^= 1;
if constexpr (DO_PROFILE) profiler.stop(TAG_TMA_WAIT, should_profile);
if constexpr (DO_PROFILE) profiler.start(should_profile);
for (int i = 0; i < NUM_MMA_ITERS; i++) {
uint64_t sdesc_a = make_matrix_descriptor_sf(&SFA[current_stage_idx * SF_bytes + i * SF_TILE_SIZE_BYTES]);
uint64_t sdesc_b1 = make_matrix_descriptor_sf(&SFB1[current_stage_idx * SF_bytes + i * SF_TILE_SIZE_BYTES]);
uint64_t sdesc_b2 = make_matrix_descriptor_sf(&SFB2[current_stage_idx * SF_bytes + i * SF_TILE_SIZE_BYTES]);
tcgen05_cp(get_tmem_addr(sfa_stage_base, 0, i * 4), sdesc_a);
tcgen05_cp(get_tmem_addr(sfb1_stage_base, 0, i * 4), sdesc_b1);
tcgen05_cp(get_tmem_addr(sfb2_stage_base, 0, i * 4), sdesc_b2);
}
if constexpr (DO_PROFILE) profiler.stop(TAG_SCALE_LOADING, should_profile);
tcgen05_commit_group(&mma_ready_bars[current_stage_idx]);
current_stage_idx = (current_stage_idx + 1) % NUM_STAGES;
}
} else if (warp_id == 6 && elect_sync()) {
uint32_t mma_phase[NUM_STAGES] = {};
int current_stage_idx = 0;
const int sfa_offset = (block_row % (128 / BLOCK_M) * (BLOCK_M / 32));
const int sfb1_offset = (block_col % (128 / BLOCK_N) * (BLOCK_N / 32));
for (int iter = 0; iter < num_iters; iter++) {
uint32_t sfa_stage_base = get_tmem_addr(sfa_tmem_addr, 0, current_stage_idx * SFA_COLS);
uint32_t sfb1_stage_base = get_tmem_addr(sfb1_tmem_addr, 0, current_stage_idx * SFB_COLS);
if constexpr (DO_PROFILE) profiler.start(should_profile);
wait(mma_ready_bars[current_stage_idx], mma_phase[current_stage_idx]);
if constexpr (DO_PROFILE) profiler.stop(TAG_TMA_WAIT, should_profile);
mma_phase[current_stage_idx] ^= 1;
if constexpr (DO_PROFILE) profiler.start(should_profile);
static_assert(BLOCK_K_IN_BYTES == K_LOAD_SIZE_IN_BYTES, "BLOCK_K_IN_BYTES must be equal to K_LOAD_SIZE_IN_BYTES");
constexpr int NUM_MMA_PER_TMA_TILE = K_LOAD_SIZE_IN_BYTES / MMA_K_IN_BYTES;
{
uint32_t sfa_addr = get_tmem_addr(sfa_stage_base, 0, sfa_offset);
uint32_t sfb1_addr = get_tmem_addr(sfb1_stage_base, 0, sfb1_offset);
if (iter == 0) {
tcgen05_mma_nvf4<true>(d1_tmem_addr, idesc_nvf4,
&As[current_stage_idx * A_bytes + 0 * MMA_K_IN_BYTES], &B1s[current_stage_idx * B_bytes + 0 * MMA_K_IN_BYTES],
sfa_addr, sfb1_addr);
} else {
tcgen05_mma_nvf4<false>(d1_tmem_addr, idesc_nvf4,
&As[current_stage_idx * A_bytes + 0 * MMA_K_IN_BYTES], &B1s[current_stage_idx * B_bytes + 0 * MMA_K_IN_BYTES],
sfa_addr, sfb1_addr);
}
}
#pragma unroll
for (int mma_iter = 1; mma_iter < NUM_MMA_PER_TMA_TILE; mma_iter++) {
uint32_t sfa_addr = get_tmem_addr(sfa_stage_base, 0, mma_iter * SF_NUM_COLS_PER_ITER + sfa_offset);
uint32_t sfb1_addr = get_tmem_addr(sfb1_stage_base, 0, mma_iter * SF_NUM_COLS_PER_ITER + sfb1_offset);
tcgen05_mma_nvf4<false>(d1_tmem_addr, idesc_nvf4,
&As[current_stage_idx * A_bytes + mma_iter * MMA_K_IN_BYTES], &B1s[current_stage_idx * B_bytes + mma_iter * MMA_K_IN_BYTES],
sfa_addr, sfb1_addr);
}
tcgen05_commit_group(&is_empty_bars[current_stage_idx]);
if constexpr (DO_PROFILE) profiler.stop(TAG_MMA_ISSUE, should_profile);
current_stage_idx = (current_stage_idx + 1) % NUM_STAGES;
}
tcgen05_commit_group(&last_stage_bar);
} else if (warp_id == 7 && elect_sync()) {
uint32_t mma_phase[NUM_STAGES] = {};
int current_stage_idx = 0;
const int sfa_offset = (block_row % (128 / BLOCK_M) * (BLOCK_M / 32));
const int sfb2_offset = (block_col % (128 / BLOCK_N) * (BLOCK_N / 32));
for (int iter = 0; iter < num_iters; iter++) {
uint32_t sfa_stage_base = get_tmem_addr(sfa_tmem_addr, 0, current_stage_idx * SFA_COLS);
uint32_t sfb2_stage_base = get_tmem_addr(sfb2_tmem_addr, 0, current_stage_idx * SFB_COLS);
if constexpr (DO_PROFILE) profiler.start(should_profile);
wait(mma_ready_bars[current_stage_idx], mma_phase[current_stage_idx]);
if constexpr (DO_PROFILE) profiler.stop(TAG_TMA_WAIT, should_profile);
mma_phase[current_stage_idx] ^= 1;
if constexpr (DO_PROFILE) profiler.start(should_profile);
static_assert(BLOCK_K_IN_BYTES == K_LOAD_SIZE_IN_BYTES, "BLOCK_K_IN_BYTES must be equal to K_LOAD_SIZE_IN_BYTES");
constexpr int NUM_MMA_PER_TMA_TILE = K_LOAD_SIZE_IN_BYTES / MMA_K_IN_BYTES;
{
uint32_t sfa_addr = get_tmem_addr(sfa_stage_base, 0, sfa_offset);
uint32_t sfb2_addr = get_tmem_addr(sfb2_stage_base, 0, sfb2_offset);
if (iter == 0) {
tcgen05_mma_nvf4<true>(d2_tmem_addr, idesc_nvf4,
&As[current_stage_idx * A_bytes + 0 * MMA_K_IN_BYTES], &B2s[current_stage_idx * B_bytes + 0 * MMA_K_IN_BYTES],
sfa_addr, sfb2_addr);
} else {
tcgen05_mma_nvf4<false>(d2_tmem_addr, idesc_nvf4,
&As[current_stage_idx * A_bytes + 0 * MMA_K_IN_BYTES], &B2s[current_stage_idx * B_bytes + 0 * MMA_K_IN_BYTES],
sfa_addr, sfb2_addr);
}
}
#pragma unroll
for (int mma_iter = 1; mma_iter < NUM_MMA_PER_TMA_TILE; mma_iter++) {
uint32_t sfa_addr = get_tmem_addr(sfa_stage_base, 0, mma_iter * SF_NUM_COLS_PER_ITER + sfa_offset);
uint32_t sfb2_addr = get_tmem_addr(sfb2_stage_base, 0, mma_iter * SF_NUM_COLS_PER_ITER + sfb2_offset);
tcgen05_mma_nvf4<false>(d2_tmem_addr, idesc_nvf4,
&As[current_stage_idx * A_bytes + mma_iter * MMA_K_IN_BYTES], &B2s[current_stage_idx * B_bytes + mma_iter * MMA_K_IN_BYTES],
sfa_addr, sfb2_addr);
}
tcgen05_commit_group(&is_empty_bars[current_stage_idx]);
if constexpr (DO_PROFILE) profiler.stop(TAG_MMA_ISSUE, should_profile);
current_stage_idx = (current_stage_idx + 1) % NUM_STAGES;
}
tcgen05_commit_group(&last_stage_bar);
} else if (warp_id < 4) {
wait(last_stage_bar, 0);
asm volatile("tcgen05.fence::after_thread_sync;");
if constexpr (DO_PROFILE) profiler.start(should_profile);
if constexpr (USE_TMA_EPILOGUE) {
constexpr int kSwizzleCDMode = 128;
constexpr int kNumBankGroupBytes = 16;
constexpr int kNumElemsPerBankGroup = kNumBankGroupBytes / sizeof(half); // 8 elements
constexpr int kNumStores = BLOCK_N / STORE_BLOCK_N;
const int lane_id = threadIdx.x % WARP_SIZE;
int row = threadIdx.x;
uint8_t* cs_base = reinterpret_cast<uint8_t*>(&Cs[0]);
int row_byte_offset = row * STORE_BLOCK_N * sizeof(half);
uint32_t row_smem_addr = cvta_shared(cs_base + row_byte_offset);
int swizzle_xor = (row_smem_addr >> 7) & 0x7;
for (int chunk = 0; chunk < kNumStores; ++chunk) {
int chunk_smem_offset_base = chunk * (BLOCK_M * STORE_BLOCK_N * sizeof(half));
int chunk_col_offset = chunk * STORE_BLOCK_N;
float d1_buf[2][8];
float d2_buf[2][8];
int stage = 0;
constexpr int num_iterations = STORE_BLOCK_N / kNumElemsPerBankGroup;
uint32_t d1_base_addr_first = get_tmem_addr(d1_tmem_addr, warp_id * 32, chunk_col_offset);
uint32_t d2_base_addr_first = get_tmem_addr(d2_tmem_addr, warp_id * 32, chunk_col_offset);
tcgen05_ld_8(d1_buf[stage], d1_base_addr_first);
tcgen05_ld_8(d2_buf[stage], d2_base_addr_first);
uint32_t d1_base_addr_next = get_tmem_addr(d1_tmem_addr, warp_id * 32, chunk_col_offset + kNumElemsPerBankGroup);
uint32_t d2_base_addr_next = get_tmem_addr(d2_tmem_addr, warp_id * 32, chunk_col_offset + kNumElemsPerBankGroup);
for (int iter = 0; iter < num_iterations; ++iter) {
int col_in_chunk = iter * kNumElemsPerBankGroup;
int bank_group_index = col_in_chunk / kNumElemsPerBankGroup;
asm volatile("tcgen05.wait::ld.sync.aligned;");
if (iter + 1 < num_iterations) {
int next_stage = stage ^ 1;
tcgen05_ld_8(d1_buf[next_stage], d1_base_addr_next);
tcgen05_ld_8(d2_buf[next_stage], d2_base_addr_next);
d1_base_addr_next = get_tmem_addr(d1_tmem_addr, warp_id * 32, chunk_col_offset + (iter + 2) * kNumElemsPerBankGroup);
d2_base_addr_next = get_tmem_addr(d2_tmem_addr, warp_id * 32, chunk_col_offset + (iter + 2) * kNumElemsPerBankGroup);
}
half2 packed[4];
#pragma unroll 2
for (int i = 0; i < 4; i++) {
float2 d1_pair = {d1_buf[stage][i*2], d1_buf[stage][i*2+1]};
float2 d2_pair = {d2_buf[stage][i*2], d2_buf[stage][i*2+1]};
// Vectorized SiLU
float2 sigmoid = {1.0f / (1.0f + __expf(-d1_pair.x)),
1.0f / (1.0f + __expf(-d1_pair.y))};
float2 result = {d1_pair.x * sigmoid.x * d2_pair.x,
d1_pair.y * sigmoid.y * d2_pair.y};
packed[i] = __float22half2_rn(result);
}
int col = bank_group_index ^ swizzle_xor;
auto smem_ptr = cs_base + row_byte_offset + col * kNumBankGroupBytes + chunk_smem_offset_base;
st_shared(smem_ptr,
*reinterpret_cast<uint32_t*>(&packed[0]),
*reinterpret_cast<uint32_t*>(&packed[1]),
*reinterpret_cast<uint32_t*>(&packed[2]),
*reinterpret_cast<uint32_t*>(&packed[3]));
stage ^= 1;
}
ptx::fence_proxy_async(ptx::space_shared);
named_barrier_sync(0, 128);
if (warp_id == 0 && elect_sync()) {
int32_t c_tensor_coords[2] = {block_col * BLOCK_N + chunk_col_offset, block_row * BLOCK_M};
ptx::cp_async_bulk_tensor(ptx::space_global, ptx::space_shared, &c_tensor_map, c_tensor_coords, Cs + chunk * (BLOCK_M * STORE_BLOCK_N));
ptx::cp_async_bulk_commit_group();
}
}
if (warp_id == 0 && elect_sync()) {
ptx::cp_async_bulk_wait_group_read(ptx::n32_t<0>());
}
} else {
const int lane_id = threadIdx.x % WARP_SIZE;
float d1_results[2][BLOCK_N / 2];
float d2_results[2][BLOCK_N / 2];
uint32_t d1_base_addr = get_tmem_addr(d1_tmem_addr, warp_id * 32, 0);
uint32_t d2_base_addr = get_tmem_addr(d2_tmem_addr, warp_id * 32, 0);
tcgen05_ld<BLOCK_N>(d1_results[0], d1_base_addr);
tcgen05_ld<BLOCK_N>(d2_results[0], d2_base_addr);
for (int batch = 0; batch < 2; batch++) {
asm volatile("tcgen05.wait::ld.sync.aligned;");
if (batch < 1) {
uint32_t d1_next_addr = get_tmem_addr(d1_tmem_addr, warp_id * 32 + 16, 0);
uint32_t d2_next_addr = get_tmem_addr(d2_tmem_addr, warp_id * 32 + 16, 0);
tcgen05_ld<BLOCK_N>(d1_results[1], d1_next_addr);
tcgen05_ld<BLOCK_N>(d2_results[1], d2_next_addr);
}
const int base_row = warp_id * 32 + batch * 16 + lane_id / 4;
const int col_base = (lane_id % 4) * 2;
#pragma unroll
for (int i = 0; i < BLOCK_N / 8; i++) {
const int col = i * 8 + col_base;
const int idx = i * 4;
float a = silu(d1_results[batch][idx + 0]) * d2_results[batch][idx + 0];
float b = silu(d1_results[batch][idx + 1]) * d2_results[batch][idx + 1];
float c = silu(d1_results[batch][idx + 2]) * d2_results[batch][idx + 2];
float d = silu(d1_results[batch][idx + 3]) * d2_results[batch][idx + 3];
*reinterpret_cast<half2*>(&C[(base_row + 0) * N + col]) = __float22half2_rn({a, b});
*reinterpret_cast<half2*>(&C[(base_row + 8) * N + col]) = __float22half2_rn({c, d});
}
}
}
named_barrier_sync(0, 128);
if constexpr (DO_PROFILE) profiler.stop(TAG_EPILOGUE, should_profile);
if (warp_id == 1) {
deallocate_tensor_memory(tmem_addr_base, 512);
}
}
}
template<const int SMEM_WIDTH, const int SMEM_HEIGHT >
CUtensorMap create_tensor_map_ab(void* global_address, int GMEM_WIDTH, int GMEM_HEIGHT) {
CUtensorMap tensor_map{};
constexpr int rank = 3;
uint32_t elem_stride[rank] = {1, 1, 1};
uint64_t size[rank] = {K_LOAD_SIZE_IN_BYTES, static_cast<uint64_t>(GMEM_HEIGHT), static_cast<uint64_t>(GMEM_WIDTH / K_LOAD_SIZE_IN_BYTES)};
uint32_t box_size[rank] = {K_LOAD_SIZE_IN_BYTES, SMEM_HEIGHT, SMEM_WIDTH / K_LOAD_SIZE_IN_BYTES};
uint64_t stride[rank - 1] = {static_cast<uint64_t>(GMEM_WIDTH), K_LOAD_SIZE_IN_BYTES};
CUresult res = cuTensorMapEncodeTiled(
&tensor_map,
CUtensorMapDataType::CU_TENSOR_MAP_DATA_TYPE_UINT8,
rank,
(void*)global_address,
size,
stride,
box_size,
elem_stride,
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
);
return tensor_map;
}
template<const int SMEM_WIDTH, const int SMEM_HEIGHT>
CUtensorMap create_tensor_map_c(void* global_address, int M, int N) {
CUtensorMap tensor_map{};
constexpr int rank = 2;
uint32_t elem_stride[rank] = {1, 1};
uint64_t size[rank] = {static_cast<uint64_t>(N), static_cast<uint64_t>(M)};
uint32_t box_size[rank] = {SMEM_WIDTH, SMEM_HEIGHT};
uint64_t stride[1] = {static_cast<uint64_t>(N * sizeof(half))};
CUresult res = cuTensorMapEncodeTiled(
&tensor_map,
CUtensorMapDataType::CU_TENSOR_MAP_DATA_TYPE_FLOAT16,
rank,
(void*)global_address,
size,
stride,
box_size,
elem_stride,
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
);
return tensor_map;
}
template<const bool DO_PROFILE, const int BLOCK_M, const int BLOCK_N, const int BLOCK_K, const int NUM_STAGES, const bool USE_TMA_EPILOGUE>
torch::Tensor matmul_impl(torch::Tensor A, torch::Tensor B1, torch::Tensor B2, torch::Tensor SFA_permuted, torch::Tensor SFB1_permuted, torch::Tensor SFB2_permuted, torch::Tensor C, int64_t* profile_buf) {
constexpr int BLOCK_K_IN_BYTES = BLOCK_K / ELEMENTS_PER_BYTE;
int M = A.size(0);
int N = B1.size(0);
int K_BYTES = A.size(1); // This is in bytes
int K = K_BYTES * 2;
constexpr int L = 1;
CUtensorMap a_tensor_map = create_tensor_map_ab<BLOCK_K_IN_BYTES, BLOCK_M>(
A.data_ptr(), K_BYTES, M); // M rows, K_BYTES cols (bytes)
CUtensorMap b1_tensor_map = create_tensor_map_ab<BLOCK_K_IN_BYTES, BLOCK_N>(
B1.data_ptr(), K_BYTES, N);
CUtensorMap b2_tensor_map = create_tensor_map_ab<BLOCK_K_IN_BYTES, BLOCK_N>(
B2.data_ptr(), K_BYTES, N);
CUtensorMap c_tensor_map = create_tensor_map_c<STORE_BLOCK_N, BLOCK_M>(C.data_ptr(), M, N);
const int NUM_BLOCKS = cdiv(N, BLOCK_N) * cdiv(M, BLOCK_M);
dim3 gridDim(NUM_BLOCKS);
dim3 blockDim(NUM_THREADS);
auto kernel = cuda_kernel<BLOCK_M, BLOCK_N, BLOCK_K, NUM_STAGES, USE_TMA_EPILOGUE, DO_PROFILE>;
constexpr int A_bytes = BLOCK_M * BLOCK_K_IN_BYTES;
constexpr int B_bytes = BLOCK_N * BLOCK_K_IN_BYTES;
constexpr int NUM_MMA_ITERS = BLOCK_K / MMA_K;
constexpr int SF_bytes = SF_TILE_SIZE_BYTES * NUM_MMA_ITERS;
constexpr int stage_size = A_bytes + B_bytes * 2 + SF_bytes * 3;
constexpr int shared_size = stage_size * NUM_STAGES;
cudaFuncSetAttribute(kernel, cudaFuncAttributeMaxDynamicSharedMemorySize, shared_size);
kernel<<<gridDim, blockDim, shared_size>>>(
M,
N,
K, // K in elements
a_tensor_map,
b1_tensor_map,
b2_tensor_map,
c_tensor_map,
reinterpret_cast<char*>(SFA_permuted.data_ptr()),
reinterpret_cast<char*>(SFB1_permuted.data_ptr()),
reinterpret_cast<char*>(SFB2_permuted.data_ptr()),
reinterpret_cast<half*>(C.data_ptr()),
profile_buf
);
return C;
}
template<const bool DO_PROFILE>
torch::Tensor matmul_impl(torch::Tensor A, torch::Tensor B1, torch::Tensor B2, torch::Tensor SFA_permuted, torch::Tensor SFB1_permuted, torch::Tensor SFB2_permuted, torch::Tensor C, int64_t* profile_buf) {
int M = A.size(0);
if (M == 256) {
return matmul_impl<DO_PROFILE, 128, 64, 256, 5, true>(A, B1, B2, SFA_permuted, SFB1_permuted, SFB2_permuted, C, profile_buf);
} else {
return matmul_impl<DO_PROFILE, 128, 128, 256, 4, true>(A, B1, B2, SFA_permuted, SFB1_permuted, SFB2_permuted, C, profile_buf);
}
}
torch::Tensor cuda_entry(torch::Tensor A, torch::Tensor B1, torch::Tensor B2, torch::Tensor SFA_permuted, torch::Tensor SFB1_permuted, torch::Tensor SFB2_permuted, torch::Tensor C) {
return matmul_impl<false>(A, B1, B2, SFA_permuted, SFB1_permuted, SFB2_permuted, C, nullptr);
}
torch::Tensor cuda_entry_with_profile(torch::Tensor A, torch::Tensor B1, torch::Tensor B2, torch::Tensor SFA_permuted, torch::Tensor SFB1_permuted, torch::Tensor SFB2_permuted, torch::Tensor C, torch::Tensor profile_buf) {
return matmul_impl<true>(A, B1, B2, SFA_permuted, SFB1_permuted, SFB2_permuted, C, reinterpret_cast<int64_t*>(profile_buf.data_ptr()));
}
"""
cuda_module = load_inline(
name="cuda_kernel",
cpp_sources=cpp_source,
cuda_sources=CUDA_SRC,
functions=["cuda_entry"],
extra_cuda_cflags=[
"-O3",
"-lineinfo",
"-Xptxas=-v",
"-gencode=arch=compute_100a,code=sm_100a",
"--use_fast_math",
],
extra_ldflags=["-lcuda"], # for cuTensorMapEncodeTiled() used by TMA
verbose=True,
)
import torch
def custom_kernel(data: input_t) -> output_t:
a, b1, b2, _, _, _, sfa_permuted, sfb1_permuted, sfb2_permuted, c = data
out = cuda_module.cuda_entry(a, b1, b2, sfa_permuted, sfb1_permuted, sfb2_permuted, c)
return outscrolls · 1023 lines total
Source code from GPU Mode and the KernelBot dataset · June 9 Researcher Reciprocity License v1.0
Changes from previous submission
Against this author's previous submission submission 341566.
⋯ 204 unchanged lines: "memory");}+ template<uint64_t CACHE_POLICY>+ __device__ __forceinline__ void cp_async_bulk_tensor_gmem2smem_multicast(void* dst, const CUtensorMap* tmap, const int32_t* coords, uint64_t* mbar) {+ uint32_t dst_addr = static_cast<uint32_t>(__cvta_generic_to_shared(dst));+ uint32_t mbar_addr = static_cast<uint32_t>(__cvta_generic_to_shared(mbar));+ asm volatile("cp.async.bulk.tensor.3d.shared::cluster.global.mbarrier::complete_tx::bytes.cta_group::1.L2::cache_hint "+ "[%0], [%1, {%2, %3, %4}], [%5], %6;"+ :: "r"(dst_addr), "l"(tmap), "r"(coords[0]), "r"(coords[1]), "r"(coords[2]), "r"(mbar_addr), "l"(CACHE_POLICY)+ : "memory");+ }+__device__ __forceinline__ void st_shared(const void* ptr, uint32_t x, uint32_t y, uint32_t z, uint32_t w) {asm volatile("st.shared.v4.u32 [%0], {%1, %2, %3, %4};" :: "l"(__cvta_generic_to_shared(ptr)), "r"(x), "r"(y), "r"(z), "r"(w));}⋯ 181 unchanged lines// nvf4 MMA with collector buffer lastuse - uses cached A from collector and discardstemplate <bool init>- __device__ __forceinline__ void tcgen05_mma_nvf4_use(+ __device__ __forceinline__ void tcgen05_mma_nvf4_lastuse(uint32_t d_tmem_addr,uint32_t i_desc,void* sA,⋯ 9 unchanged lines"{\n"".reg .pred p;\n""setp.eq.u32 p, 1, 0;\n" // p = False (initialize)- "tcgen05.mma.cta_group::1.kind::mxf4nvf4.block_scale.block16.collector::a::use "+ "tcgen05.mma.cta_group::1.kind::mxf4nvf4.block_scale.block16.collector::a::lastuse ""[%0], %1, %2, %3, [%4], [%5], p;\n""}\n":: "r"(d_tmem_addr), "l"(a_desc), "l"(b_desc), "r"(i_desc),⋯ 4 unchanged lines"{\n"".reg .pred p;\n""setp.eq.u32 p, 1, 1;\n" // p = True (accumulate)- "tcgen05.mma.cta_group::1.kind::mxf4nvf4.block_scale.block16.collector::a::use "+ "tcgen05.mma.cta_group::1.kind::mxf4nvf4.block_scale.block16.collector::a::lastuse ""[%0], %1, %2, %3, [%4], [%5], p;\n""}\n":: "r"(d_tmem_addr), "l"(a_desc), "l"(b_desc), "r"(i_desc),⋯ 17 unchanged linestemplate<const int BLOCK_M, const int BLOCK_N, const int BLOCK_K, const int NUM_STAGES, const bool USE_TMA_EPILOGUE, const bool DO_PROFILE>__launch_bounds__(NUM_THREADS)+ __cluster_dims__(2, 1, 1)__global__ void cuda_kernel(const int M,const int N,⋯ 25 unchanged linesconstexpr uint32_t idesc_nvf4 =(1 << 7) | (1 << 10) | ((BLOCK_N >> 3) << 17) | ((BLOCK_M >> 7) << 27);+const int warp_id = threadIdx.x / WARP_SIZE;const int block_idx = blockIdx.x;- const int block_row = block_idx / (N / BLOCK_N);- const int block_col = block_idx % (N / BLOCK_N);+ uint32_t cluster_ctaid;+ asm volatile("mov.b32 %0, %%cluster_ctarank;" : "=r"(cluster_ctaid));+ // const int block_row = block_idx / (N / BLOCK_N);+ // const int block_col = block_idx % (N / BLOCK_N);+ const int cluster_idx = blockIdx.x / 2;+ const int blocks_per_row = N / BLOCK_N;+ const int cluster_row = (cluster_idx / blocks_per_row) * 2;+ const int cluster_col = cluster_idx % blocks_per_row;+ const int block_row = cluster_row + cluster_ctaid;+ const int block_col = cluster_col;+IntraKernelProfiler profiler;bool should_profile = false;⋯ 28 unchanged linesallocate_tensor_memory(&tmem_addr_base_shared, 512);}- // Each stage uses SFA_COLS + SFB_COLS * 2 columns- constexpr int SF_STAGE_STRIDE = SFA_COLS + SFB_COLS * 2;+ // Tensor memory layout: SFA * NUM_STAGES, SFB1 * NUM_STAGES, SFB2 * NUM_STAGES+ constexpr int SF_TOTAL_COLS = SFA_COLS * NUM_STAGES + SFB_COLS * NUM_STAGES * 2;__shared__ __align__(8) uint64_t is_empty_bars[NUM_STAGES];__shared__ __align__(8) uint64_t mma_ready_bars[NUM_STAGES];⋯ 1 unchanged lines__shared__ __align__(8) uint64_t last_stage_bar;if (threadIdx.x < NUM_STAGES) {- init_barrier(&is_empty_bars[threadIdx.x], 1);- init_barrier(&mma_ready_bars[threadIdx.x], 3);+ init_barrier(&is_empty_bars[threadIdx.x], 2);+ init_barrier(&mma_ready_bars[threadIdx.x], 2);init_barrier(&sf_tma_ready_bars[threadIdx.x], 1);}if (threadIdx.x == 0) {- init_barrier(&last_stage_bar, 1);+ init_barrier(&last_stage_bar, 2);}__syncthreads();const uint32_t tmem_addr_base = tmem_addr_base_shared;const uint32_t sfa_tmem_addr = get_tmem_addr(tmem_addr_base, 0, 0);- const uint32_t sfb1_tmem_addr = get_tmem_addr(sfa_tmem_addr, 0, SFA_COLS);- const uint32_t sfb2_tmem_addr = get_tmem_addr(sfb1_tmem_addr, 0, SFB_COLS);- const uint32_t d1_tmem_addr = get_tmem_addr(tmem_addr_base, 0, NUM_STAGES * SF_STAGE_STRIDE);+ const uint32_t sfb1_tmem_addr = get_tmem_addr(sfa_tmem_addr, 0, SFA_COLS * NUM_STAGES);+ const uint32_t sfb2_tmem_addr = get_tmem_addr(sfb1_tmem_addr, 0, SFB_COLS * NUM_STAGES);+ const uint32_t d1_tmem_addr = get_tmem_addr(sfb2_tmem_addr, 0, SFB_COLS * NUM_STAGES);const uint32_t d2_tmem_addr = get_tmem_addr(tmem_addr_base, 0, 512 - BLOCK_N);⋯ 12 unchanged linesauto issue_tma = [&](int iter, int stage_id) {if constexpr (DO_PROFILE) profiler.start(should_profile);- expect_bytes(&sf_tma_ready_bars[stage_id], sf_tma_bytes);constexpr int NUM_K_CHUNKS = BLOCK_K_IN_BYTES / K_LOAD_SIZE_IN_BYTES;const int k_chunk_start = iter * NUM_K_CHUNKS;int32_t a_tensor_coords[3] = {0, row_start, k_chunk_start};int32_t b12_tensor_coords[3] = {0, col_start, k_chunk_start};- cp_async_bulk_gmem2smem<EVICT_LAST>(&SFA[stage_id * SF_bytes], sfa_ptr + iter * SF_bytes, SF_bytes, &sf_tma_ready_bars[stage_id]);- cp_async_bulk_gmem2smem<EVICT_LAST>(&SFB1[stage_id * SF_bytes], sfb1_ptr + iter * SF_bytes, SF_bytes, &sf_tma_ready_bars[stage_id]);- cp_async_bulk_gmem2smem<EVICT_LAST>(&SFB2[stage_id * SF_bytes], sfb2_ptr + iter * SF_bytes, SF_bytes, &sf_tma_ready_bars[stage_id]);-- expect_bytes(&mma_ready_bars[stage_id], matrix_bytes);+ cp_async_bulk_gmem2smem<EVICT_FIRST>(&SFA[stage_id * SF_bytes], sfa_ptr + iter * SF_bytes, SF_bytes, &sf_tma_ready_bars[stage_id]);+ cp_async_bulk_gmem2smem<EVICT_FIRST>(&SFB1[stage_id * SF_bytes], sfb1_ptr + iter * SF_bytes, SF_bytes, &sf_tma_ready_bars[stage_id]);+ cp_async_bulk_gmem2smem<EVICT_FIRST>(&SFB2[stage_id * SF_bytes], sfb2_ptr + iter * SF_bytes, SF_bytes, &sf_tma_ready_bars[stage_id]);+ expect_bytes(&sf_tma_ready_bars[stage_id], sf_tma_bytes);- cp_async_bulk_tensor_gmem2smem<EVICT_NORMAL>(&As[stage_id * A_bytes], &a_tensor_map, a_tensor_coords, &mma_ready_bars[stage_id]);+ cp_async_bulk_tensor_gmem2smem<EVICT_LAST>(&As[stage_id * A_bytes], &a_tensor_map, a_tensor_coords, &mma_ready_bars[stage_id]);cp_async_bulk_tensor_gmem2smem<EVICT_NORMAL>(&B1s[stage_id * B_bytes], &b1_tensor_map, b12_tensor_coords, &mma_ready_bars[stage_id]);cp_async_bulk_tensor_gmem2smem<EVICT_NORMAL>(&B2s[stage_id * B_bytes], &b2_tensor_map, b12_tensor_coords, &mma_ready_bars[stage_id]);+ expect_bytes(&mma_ready_bars[stage_id], matrix_bytes);+if constexpr (DO_PROFILE) profiler.stop(TAG_TMA_ISSUE, should_profile);};⋯ 18 unchanged linesint current_stage_idx = 0;for (int iter = 0; iter < num_iters; iter++) {- uint32_t sfa_stage_base = get_tmem_addr(sfa_tmem_addr, 0, current_stage_idx * SF_STAGE_STRIDE);- // uint32_t sfb1_stage_base = get_tmem_addr(sfb1_tmem_addr, 0, current_stage_idx * SF_STAGE_STRIDE);- uint32_t sfb2_stage_base = get_tmem_addr(sfb2_tmem_addr, 0, current_stage_idx * SF_STAGE_STRIDE);+ uint32_t sfa_stage_base = get_tmem_addr(sfa_tmem_addr, 0, current_stage_idx * SFA_COLS);+ uint32_t sfb1_stage_base = get_tmem_addr(sfb1_tmem_addr, 0, current_stage_idx * SFB_COLS);+ uint32_t sfb2_stage_base = get_tmem_addr(sfb2_tmem_addr, 0, current_stage_idx * SFB_COLS);if constexpr (DO_PROFILE) profiler.start(should_profile);wait(sf_tma_ready_bars[current_stage_idx], sf_phase[current_stage_idx]);⋯ 3 unchanged linesfor (int i = 0; i < NUM_MMA_ITERS; i++) {uint64_t sdesc_a = make_matrix_descriptor_sf(&SFA[current_stage_idx * SF_bytes + i * SF_TILE_SIZE_BYTES]);- // uint64_t sdesc_b1 = make_matrix_descriptor_sf(&SFB1[current_stage_idx * SF_bytes + i * SF_TILE_SIZE_BYTES]);+ uint64_t sdesc_b1 = make_matrix_descriptor_sf(&SFB1[current_stage_idx * SF_bytes + i * SF_TILE_SIZE_BYTES]);uint64_t sdesc_b2 = make_matrix_descriptor_sf(&SFB2[current_stage_idx * SF_bytes + i * SF_TILE_SIZE_BYTES]);tcgen05_cp(get_tmem_addr(sfa_stage_base, 0, i * 4), sdesc_a);- // tcgen05_cp(get_tmem_addr(sfb1_stage_base, 0, i * 4), sdesc_b1);+ tcgen05_cp(get_tmem_addr(sfb1_stage_base, 0, i * 4), sdesc_b1);tcgen05_cp(get_tmem_addr(sfb2_stage_base, 0, i * 4), sdesc_b2);}if constexpr (DO_PROFILE) profiler.stop(TAG_SCALE_LOADING, should_profile);⋯ 2 unchanged linescurrent_stage_idx = (current_stage_idx + 1) % NUM_STAGES;}} else if (warp_id == 6 && elect_sync()) {- uint32_t sf_phase[NUM_STAGES] = {};+ uint32_t mma_phase[NUM_STAGES] = {};int current_stage_idx = 0;++ const int sfa_offset = (block_row % (128 / BLOCK_M) * (BLOCK_M / 32));+ const int sfb1_offset = (block_col % (128 / BLOCK_N) * (BLOCK_N / 32));for (int iter = 0; iter < num_iters; iter++) {+ uint32_t sfa_stage_base = get_tmem_addr(sfa_tmem_addr, 0, current_stage_idx * SFA_COLS);+ uint32_t sfb1_stage_base = get_tmem_addr(sfb1_tmem_addr, 0, current_stage_idx * SFB_COLS);- // uint32_t sfa_stage_base = get_tmem_addr(sfa_tmem_addr, 0, current_stage_idx * SF_STAGE_STRIDE);- uint32_t sfb1_stage_base = get_tmem_addr(sfb1_tmem_addr, 0, current_stage_idx * SF_STAGE_STRIDE);- // uint32_t sfb2_stage_base = get_tmem_addr(sfb2_tmem_addr, 0, current_stage_idx * SF_STAGE_STRIDE);-if constexpr (DO_PROFILE) profiler.start(should_profile);- wait(sf_tma_ready_bars[current_stage_idx], sf_phase[current_stage_idx]);- sf_phase[current_stage_idx] ^= 1;+ wait(mma_ready_bars[current_stage_idx], mma_phase[current_stage_idx]);if constexpr (DO_PROFILE) profiler.stop(TAG_TMA_WAIT, should_profile);+ mma_phase[current_stage_idx] ^= 1;+if constexpr (DO_PROFILE) profiler.start(should_profile);-- for (int i = 0; i < NUM_MMA_ITERS; i++) {- // uint64_t sdesc_a = make_matrix_descriptor_sf(&SFA[current_stage_idx * SF_bytes + i * SF_TILE_SIZE_BYTES]);- uint64_t sdesc_b1 = make_matrix_descriptor_sf(&SFB1[current_stage_idx * SF_bytes + i * SF_TILE_SIZE_BYTES]);- // uint64_t sdesc_b2 = make_matrix_descriptor_sf(&SFB2[current_stage_idx * SF_bytes + i * SF_TILE_SIZE_BYTES]);- // tcgen05_cp(get_tmem_addr(sfa_stage_base, 0, i * 4), sdesc_a);- tcgen05_cp(get_tmem_addr(sfb1_stage_base, 0, i * 4), sdesc_b1);- // tcgen05_cp(get_tmem_addr(sfb2_stage_base, 0, i * 4), sdesc_b2);++ static_assert(BLOCK_K_IN_BYTES == K_LOAD_SIZE_IN_BYTES, "BLOCK_K_IN_BYTES must be equal to K_LOAD_SIZE_IN_BYTES");++ constexpr int NUM_MMA_PER_TMA_TILE = K_LOAD_SIZE_IN_BYTES / MMA_K_IN_BYTES;++ {+ uint32_t sfa_addr = get_tmem_addr(sfa_stage_base, 0, sfa_offset);+ uint32_t sfb1_addr = get_tmem_addr(sfb1_stage_base, 0, sfb1_offset);++ if (iter == 0) {+ tcgen05_mma_nvf4<true>(d1_tmem_addr, idesc_nvf4,+ &As[current_stage_idx * A_bytes + 0 * MMA_K_IN_BYTES], &B1s[current_stage_idx * B_bytes + 0 * MMA_K_IN_BYTES],+ sfa_addr, sfb1_addr);+ } else {+ tcgen05_mma_nvf4<false>(d1_tmem_addr, idesc_nvf4,+ &As[current_stage_idx * A_bytes + 0 * MMA_K_IN_BYTES], &B1s[current_stage_idx * B_bytes + 0 * MMA_K_IN_BYTES],+ sfa_addr, sfb1_addr);+ }}- if constexpr (DO_PROFILE) profiler.stop(TAG_SCALE_LOADING, should_profile);-- tcgen05_commit_group(&mma_ready_bars[current_stage_idx]);++ #pragma unroll+ for (int mma_iter = 1; mma_iter < NUM_MMA_PER_TMA_TILE; mma_iter++) {+ uint32_t sfa_addr = get_tmem_addr(sfa_stage_base, 0, mma_iter * SF_NUM_COLS_PER_ITER + sfa_offset);+ uint32_t sfb1_addr = get_tmem_addr(sfb1_stage_base, 0, mma_iter * SF_NUM_COLS_PER_ITER + sfb1_offset);+ tcgen05_mma_nvf4<false>(d1_tmem_addr, idesc_nvf4,+ &As[current_stage_idx * A_bytes + mma_iter * MMA_K_IN_BYTES], &B1s[current_stage_idx * B_bytes + mma_iter * MMA_K_IN_BYTES],+ sfa_addr, sfb1_addr);+ }+ tcgen05_commit_group(&is_empty_bars[current_stage_idx]);+ if constexpr (DO_PROFILE) profiler.stop(TAG_MMA_ISSUE, should_profile);current_stage_idx = (current_stage_idx + 1) % NUM_STAGES;}+ tcgen05_commit_group(&last_stage_bar);} else if (warp_id == 7 && elect_sync()) {uint32_t mma_phase[NUM_STAGES] = {};int current_stage_idx = 0;const int sfa_offset = (block_row % (128 / BLOCK_M) * (BLOCK_M / 32));- const int sfb1_offset = (block_col % (128 / BLOCK_N) * (BLOCK_N / 32));const int sfb2_offset = (block_col % (128 / BLOCK_N) * (BLOCK_N / 32));for (int iter = 0; iter < num_iters; iter++) {- uint32_t sfa_stage_base = get_tmem_addr(sfa_tmem_addr, 0, current_stage_idx * SF_STAGE_STRIDE);- uint32_t sfb1_stage_base = get_tmem_addr(sfb1_tmem_addr, 0, current_stage_idx * SF_STAGE_STRIDE);- uint32_t sfb2_stage_base = get_tmem_addr(sfb2_tmem_addr, 0, current_stage_idx * SF_STAGE_STRIDE);+ uint32_t sfa_stage_base = get_tmem_addr(sfa_tmem_addr, 0, current_stage_idx * SFA_COLS);+ uint32_t sfb2_stage_base = get_tmem_addr(sfb2_tmem_addr, 0, current_stage_idx * SFB_COLS);-if constexpr (DO_PROFILE) profiler.start(should_profile);wait(mma_ready_bars[current_stage_idx], mma_phase[current_stage_idx]);if constexpr (DO_PROFILE) profiler.stop(TAG_TMA_WAIT, should_profile);mma_phase[current_stage_idx] ^= 1;-if constexpr (DO_PROFILE) profiler.start(should_profile);⋯ 3 unchanged lines{uint32_t sfa_addr = get_tmem_addr(sfa_stage_base, 0, sfa_offset);- uint32_t sfb1_addr = get_tmem_addr(sfb1_stage_base, 0, sfb1_offset);uint32_t sfb2_addr = get_tmem_addr(sfb2_stage_base, 0, sfb2_offset);if (iter == 0) {- tcgen05_mma_nvf4_fill<true>(d1_tmem_addr, idesc_nvf4,- &As[current_stage_idx * A_bytes + 0 * MMA_K_IN_BYTES], &B1s[current_stage_idx * B_bytes + 0 * MMA_K_IN_BYTES],- sfa_addr, sfb1_addr);- tcgen05_mma_nvf4_use<true>(d2_tmem_addr, idesc_nvf4,+ tcgen05_mma_nvf4<true>(d2_tmem_addr, idesc_nvf4,&As[current_stage_idx * A_bytes + 0 * MMA_K_IN_BYTES], &B2s[current_stage_idx * B_bytes + 0 * MMA_K_IN_BYTES],sfa_addr, sfb2_addr);} else {- tcgen05_mma_nvf4_fill<false>(d1_tmem_addr, idesc_nvf4,- &As[current_stage_idx * A_bytes + 0 * MMA_K_IN_BYTES], &B1s[current_stage_idx * B_bytes + 0 * MMA_K_IN_BYTES],- sfa_addr, sfb1_addr);- tcgen05_mma_nvf4_use<false>(d2_tmem_addr, idesc_nvf4,+ tcgen05_mma_nvf4<false>(d2_tmem_addr, idesc_nvf4,&As[current_stage_idx * A_bytes + 0 * MMA_K_IN_BYTES], &B2s[current_stage_idx * B_bytes + 0 * MMA_K_IN_BYTES],sfa_addr, sfb2_addr);}⋯ 2 unchanged lines#pragma unrollfor (int mma_iter = 1; mma_iter < NUM_MMA_PER_TMA_TILE; mma_iter++) {uint32_t sfa_addr = get_tmem_addr(sfa_stage_base, 0, mma_iter * SF_NUM_COLS_PER_ITER + sfa_offset);- uint32_t sfb1_addr = get_tmem_addr(sfb1_stage_base, 0, mma_iter * SF_NUM_COLS_PER_ITER + sfb1_offset);uint32_t sfb2_addr = get_tmem_addr(sfb2_stage_base, 0, mma_iter * SF_NUM_COLS_PER_ITER + sfb2_offset);- tcgen05_mma_nvf4_fill<false>(d1_tmem_addr, idesc_nvf4,- &As[current_stage_idx * A_bytes + mma_iter * MMA_K_IN_BYTES], &B1s[current_stage_idx * B_bytes + mma_iter * MMA_K_IN_BYTES],- sfa_addr, sfb1_addr);- tcgen05_mma_nvf4_use<false>(d2_tmem_addr, idesc_nvf4,+ tcgen05_mma_nvf4<false>(d2_tmem_addr, idesc_nvf4,&As[current_stage_idx * A_bytes + mma_iter * MMA_K_IN_BYTES], &B2s[current_stage_idx * B_bytes + mma_iter * MMA_K_IN_BYTES],sfa_addr, sfb2_addr);}⋯ 55 unchanged lineshalf2 packed[4];+ #pragma unroll 2for (int i = 0; i < 4; i++) {float2 d1_pair = {d1_buf[stage][i*2], d1_buf[stage][i*2+1]};float2 d2_pair = {d2_buf[stage][i*2], d2_buf[stage][i*2+1]};⋯ 219 unchanged linesextra_ldflags=["-lcuda"], # for cuTensorMapEncodeTiled() used by TMAverbose=True,)+ import torchdef custom_kernel(data: input_t) -> output_t:a, b1, b2, _, _, _, sfa_permuted, sfb1_permuted, sfb2_permuted, c = data
scrolls · 317 diff lines total
Best evidence level for this revision: reported
JSON