submission 109661
v0i0 · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 758 lines, June 9 Researcher Reciprocity License v1.0.
submission-27-20-58-58.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-nvfp4-gemv-109661?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:4efb926fbfe1507103e5c99c6f9427bee23a8dc2e28c10bdb8d3605fb36cde98
license declaredunknown
license concludedunknown
authorsv0i0
imported2026-08-15
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
async-copy
__device__ void cp_async_wait_group() {fp8
const __nv_fp8_e4m3 * __restrict__ ptr_sfa,mbarrier
again: mbarrier.try_wait.parity.shared::cta.b64 done, [%0], %1, 0x100000;shared-memory
__shared__ clock_t timing[3][128];split-k
"SPLIT_K": 1,tcgen05
asm volatile("tcgen05.alloc.cta_group::1.sync.aligned.shared::cta.b32 [%0], 512;" :: "r"((unsigned int) __cvta_generic_to_shared(&tmem_addr)));tile-k = 256
const int BLOCK_K = 256;tile-m = 128
const int BLOCK_M = 128;tma
__grid_constant__ const CUtensorMap desc_avector-width = uint4
using vec_type = uint4;Kernel source
submission-27-20-58-58.py758 lines
import torch
from task import input_t, output_t
from torch.utils.cpp_extension import load_inline
from triton.testing import do_bench
import itertools
def kernel_key(*args):
return tuple([(tuple(t.shape), tuple(t.stride())) for t in args])
TUNABLES = {
# "USE_PDL": [0, 1],
# "STAGE_COUNT": [4, 6, 7, 8, 9],
# "CLUSTER_M": [None, 1, 2, 4],
# "LOCAL_M": [1, 2],
# "SPLIT_K": [1, 2],
}
def heuristic(a, b, sfa, sfb, c):
defines = {
"LOCAL_M": 1,
"USE_PDL": 1,
"CLUSTER_M": None,
"SPLIT_K": 1,
"STAGE_COUNT": 6,
}
if a.shape[0] // 128 * a.shape[2] < 64 and a.shape[1] % 512 == 0:
defines["SPLIT_K"] = 2
if a.shape[0] // 128 * a.shape[2] > 128 and a.shape[0] % 256 == 0:
defines["LOCAL_M"] = 2
return defines
def tune_all(a, b, sfa, sfb, c):
results = []
baseline = heuristic(a, b, sfa, sfb, c)
for idx, values in enumerate(itertools.product(*TUNABLES.values())):
tunables = baseline | dict(zip(TUNABLES.keys(), values))
try:
kernel = compile_kernel(kernel_key(a, b, sfa, sfb, c), tunables)
time = do_bench(lambda: kernel.run(a, b, sfa, sfb, c), warmup=5, rep=10)
print(f"{tunables} {time * 1000:.2f}")
results.append((time, idx, tunables))
except Exception as e:
print(tunables, e)
return min(results)[-1]
def tune_coordinate_descent(a, b, sfa, sfb, c):
results = {}
current = heuristic(a, b, sfa, sfb, c)
directions = list(TUNABLES.values())
def compile_kernel(key, tunables):
# print(key, tunables)
(
(a_shape, a_stride),
(_, b_stride),
(_, sfa_stride),
(_, sfb_stride),
(_, c_stride),
) = key
defines = {
"STATIC_K": 2 * a_shape[1],
"STRIDE_A_M": a_stride[0],
"STRIDE_A_L": a_stride[2],
"STRIDE_B_L": b_stride[2],
"STRIDE_SFA_M": sfa_stride[2],
"STRIDE_SFA_K": sfa_stride[4],
"STRIDE_SFA_L": sfa_stride[5],
"STRIDE_SFB_K": sfb_stride[4],
"STRIDE_SFB_L": sfb_stride[5],
"STRIDE_C_L": c_stride[2],
} | tunables
defines = [f"-D{k}={v}" for k, v in defines.items() if v is not None]
return load_inline(
cuda_sources=[
"""
#ifdef BUILD_PYTORCH
#endif
#include <cuda_fp4.h>
#include <cuda_fp8.h>
#include <cuda_fp16.h>
#include <cuda_runtime.h>
#include <cuda.h>
#include <vector>
#include <cstdint>
#define ASM(x...) asm(#x
#define ASMV(x...) asm volatile(#x
#define CLO(x...) : x)
#ifndef STATIC_K
#define STATIC_K (2*1024)
#define STRIDE_A_M 1024
#define STRIDE_A_L (1024*1024)
#define STRIDE_B_L (1024*1024)
#define STRIDE_C_L (1024*1024)
#define STRIDE_SFA_M 16384
#define STRIDE_SFA_K (4*16384)
#define STRIDE_SFA_L (8*16384)
#define STRIDE_SFB_K (4*16384)
#define STRIDE_SFB_L (8*16384)
#define LOCAL_M 1
#define SPLIT_K 1
#define STAGE_COUNT 6
#define USE_PDL 1
#endif
__device__ int semaphore[128][32] = {0};
const int BLOCK_M = 128;
const int BLOCK_K = 256;
template<int N>
__device__ void cp_async_wait_group() {
asm("cp.async.wait_group %0;" :: "n"(N) : "memory");
}
__device__ void copy_a_vec(void* dst, const void* src) {
asm("cp.async.cg.shared.global [%0], [%1], 16;" :: "r"((unsigned int)__cvta_generic_to_shared(dst)), "l"(src) : "memory");
};
__device__ void copy_b_vec(void* dst, const void* src) {
asm("cp.async.ca.shared.global [%0], [%1], 16;" :: "r"((unsigned int)__cvta_generic_to_shared(dst)), "l"(src) : "memory");
};
__device__ void copy_sf_vec(void* dst, const void* src) {
asm("cp.async.ca.shared.global [%0], [%1], 4;" :: "r"((unsigned int)__cvta_generic_to_shared(dst)), "l"(src) : "memory");
};
__device__ unsigned int make_warp_uniform(unsigned int v) {
return __shfl_sync(0xffffffff, v, 0);
}
__device__ void barrier_wait(unsigned long long* barrier, int barrier_wait_phase) {
ASMV({
.reg .pred done;
again: mbarrier.try_wait.parity.shared::cta.b64 done, [%0], %1, 0x100000;
@done bra end;
bra again;
end:
}) CLO(: "r"((unsigned int) __cvta_generic_to_shared(barrier)), "r"(barrier_wait_phase));
}
#ifdef TIMING
#define TIME_EVENT(event_id) do { \
if (is_timer) { \
timing[time_slot][current_slot] = clock(); \
tag[time_slot][current_slot] = event_id; \
current_slot += 1; \
} \
} while (0)
#else
#define TIME_EVENT(x) do {} while(0)
#endif
__launch_bounds__(256, 1)
__global__ void kernel(
const __nv_fp4x2_e2m1 * __restrict__ ptr_a,
const __nv_fp8_e4m3 * __restrict__ ptr_sfa,
const __nv_fp4x2_e2m1 * __restrict__ ptr_b,
const __nv_fp8_e4m3 * __restrict__ ptr_sfb,
__half * __restrict__ ptr_c,
__grid_constant__ const CUtensorMap desc_a
) {
#ifdef TIMING
__shared__ clock_t timing[3][128];
__shared__ int tag[3][128];
bool is_timer = threadIdx.x == 0 || threadIdx.x == 128 || threadIdx.x == (128+32);
int time_slot = threadIdx.x == 0 ? 0 : threadIdx.x == 128 ? 1 : 2;
int current_slot = 0;
TIME_EVENT(0);
#endif
// let's do the "real" blocking: 128x64 in M and K, block along M and L
ptr_a += blockIdx.y * STRIDE_A_L;
ptr_b += blockIdx.y * STRIDE_B_L;
ptr_sfa += blockIdx.y * STRIDE_SFA_L;
ptr_sfb += blockIdx.y * STRIDE_SFB_L;
ptr_c += blockIdx.y * STRIDE_C_L;
int block_offset_m = blockIdx.x / SPLIT_K;
int block_offset_k = blockIdx.x % SPLIT_K;
int k = STATIC_K / SPLIT_K;
ptr_a += block_offset_m * BLOCK_M * STRIDE_A_M * LOCAL_M;
ptr_sfa += block_offset_m * (BLOCK_M / 128) * STRIDE_SFA_M * LOCAL_M;
ptr_c += block_offset_m * BLOCK_M * LOCAL_M;
ptr_a += block_offset_k * (k / 2);
ptr_b += block_offset_k * (k / 2);
ptr_sfa += block_offset_k * (k / 16 / 4) * STRIDE_SFA_K;
ptr_sfb += block_offset_k * (k / 16 / 4) * STRIDE_SFB_K;
// have 128 threads, each computing one output element
// SF: MN, K -> addr is ((32, 4, RM), (16, 4, RK)):((16, 4, ?), (0, 1, ?))
__attribute__((aligned(128))) __shared__ __nv_fp8_e4m3 smem_sfb[STAGE_COUNT][4 * BLOCK_K / 16];
__attribute__((aligned(1024))) __shared__ __nv_fp4x2_e2m1 smem_b[STAGE_COUNT][BLOCK_K / 2];
__attribute__((aligned(128))) __shared__ __nv_fp8_e4m3 smem_sfa[STAGE_COUNT][BLOCK_M * (BLOCK_K / 16)];
// 9 * 16, 2 * 9 * 17
__attribute__((aligned(1024))) __shared__ __nv_fp4x2_e2m1 smem_a[STAGE_COUNT][BLOCK_M * (BLOCK_K / 2)];
int warp_id = threadIdx.x / 32;
int wg_id = warp_id / 4;
bool is_load = wg_id == 0;
bool is_mma = warp_id == 4;
bool is_reset = warp_id == 5;
bool is_tma = warp_id == 6;
int lane_id = threadIdx.x % 32;
__shared__ int tmem_addr;
__shared__ unsigned long long barrier;
if (warp_id == 0 && lane_id == 0) {
asm volatile("mbarrier.init.shared::cta.b64 [%0], 1;" :: "r"((unsigned int) __cvta_generic_to_shared(&barrier)));
}
int barrier_wait_phase = 0;
__shared__ unsigned long long barrier_smem_empty[STAGE_COUNT];
if (warp_id == 1 && lane_id < STAGE_COUNT) {
asm volatile("mbarrier.init.shared::cta.b64 [%0], 1;" :: "r"((unsigned int) __cvta_generic_to_shared(&barrier_smem_empty[lane_id])));
}
__shared__ unsigned long long barrier_smem_full[STAGE_COUNT];
if (warp_id == 2 && lane_id < STAGE_COUNT) {
asm volatile("mbarrier.init.shared::cta.b64 [%0], 128;" :: "r"((unsigned int) __cvta_generic_to_shared(&barrier_smem_full[lane_id])));
}
if (warp_id == 3) {
ASMV(prefetch.tensormap [%0];) CLO(: "l"(&desc_a));
}
TIME_EVENT(1);
__syncthreads();
TIME_EVENT(2);
auto mde = [](unsigned int x) { return (unsigned long long) ((x & 0x3FFFF) >> 4); };
auto make_smem_desc = [mde](void* smem_ptr, unsigned long long leading_dimension, unsigned long long stride_dimension, unsigned long long matrix_base_offset, unsigned long long is_leading_byte_address_absolute, unsigned long long swizzling_mode) {
return mde(__cvta_generic_to_shared(smem_ptr)) | (mde(leading_dimension) << 16) | (mde(stride_dimension) << 32) | (1ull << 46) | (matrix_base_offset << 49) | (is_leading_byte_address_absolute << 52) | (swizzling_mode << 61);
};
auto offset_a = [&](int m, int k) {
return (k % 16) + 16 * (((k / 16) % 8) ^ (m % 8)) + 128 * m;
};
using sf_vec_type = unsigned int;
int cutover_k = (STAGE_COUNT / LOCAL_M) * BLOCK_K;
if (cutover_k > k) cutover_k = k;
cutover_k = 0;
if (is_load) {
const int vec_size_elems = 32;
const int threads_in_k = BLOCK_K / vec_size_elems;
const int threads_in_m = BLOCK_M / threads_in_k;
using vec_type = uint4;
int ptr_a_m_offset = threadIdx.x / threads_in_k;
int ptr_a_k_offset = (threadIdx.x % threads_in_k) * sizeof(vec_type);
ptr_a += ptr_a_k_offset + ptr_a_m_offset * STRIDE_A_M;
int ptr_b_n_offset = threadIdx.x / threads_in_k;
int ptr_b_k_offset = (threadIdx.x % threads_in_k) * sizeof(vec_type);
auto smem_a_offs = reinterpret_cast<__nv_fp4x2_e2m1 (*)[BLOCK_M*BLOCK_K/2]>(&smem_a[0][offset_a(ptr_a_m_offset, ptr_a_k_offset)]);
// ptr_sfa += threadIdx.x * 16;
// auto smem_sfa_offs = reinterpret_cast<__nv_fp8_e4m3 (*)[BLOCK_M*BLOCK_K/16]>(&smem_a[0][(threadIdx.x % 32) * 16 + 128 * 4 * (threadIdx.x / 32)]);
int has_sfb = threadIdx.x * sizeof(sf_vec_type) < BLOCK_K / 16;
ptr_b += ptr_b_k_offset;
int load_stage = 0;
int load_phase = 2; // modify to not wait initially?
asm volatile("griddepcontrol.wait;" ::: "memory");
for (int block_k = 0; block_k < cutover_k; block_k += BLOCK_K) {
#pragma unroll
for (int block_m = 0; block_m < LOCAL_M; block_m++) {
TIME_EVENT(3);
if (load_phase != 2) {
barrier_wait(&barrier_smem_empty[load_stage], load_phase);
}
TIME_EVENT(4);
#pragma unroll
for (int partition_m = 0; partition_m < BLOCK_M; partition_m += threads_in_m) {
copy_a_vec(&smem_a_offs[load_stage][partition_m * 128], &ptr_a[block_k / 2 + partition_m * STRIDE_A_M + block_m * BLOCK_M * STRIDE_A_M]);
}
if (ptr_b_n_offset == 0) {
copy_b_vec(&smem_b[load_stage][ptr_b_k_offset], &ptr_b[block_k / 2]);
}
#pragma unroll
for (int rest_k = threadIdx.x / 32; rest_k < BLOCK_K / 16 / 4; rest_k += BLOCK_M / 32) {
copy_a_vec(&smem_sfa[load_stage][(threadIdx.x % 32) * 16 + 128 * 4 * rest_k], &ptr_sfa[(threadIdx.x % 32) * 16 + (rest_k + block_k / 64) * STRIDE_SFA_K + block_m * (BLOCK_M / 128) * STRIDE_SFA_M]);
}
if (has_sfb) {
int partition_k = threadIdx.x * sizeof(sf_vec_type);
copy_sf_vec(&smem_sfb[load_stage][4 * partition_k], &ptr_sfb[((block_k / 16 + partition_k) % 4) + (block_k / 64 + partition_k / (64 / 16)) * STRIDE_SFB_K]);
}
ASMV(cp.async.mbarrier.arrive.noinc.shared::cta.b64 [%0];) CLO(: "r"((unsigned int) __cvta_generic_to_shared(&barrier_smem_full[load_stage])));
load_stage += 1;
if (load_stage == STAGE_COUNT) {
load_stage = 0;
load_phase = load_phase ? 0 : 1;
}
}
}
ASMV(bar.arrive 1, 160;) CLO();
// start computing
} else if (is_mma) {
asm volatile("tcgen05.alloc.cta_group::1.sync.aligned.shared::cta.b32 [%0], 512;" :: "r"((unsigned int) __cvta_generic_to_shared(&tmem_addr)));
int read_stage = 0;
int read_phase = 0;
for (int block_k = 0; block_k < k; block_k += BLOCK_K) {
#pragma unroll
for (int block_m = 0; block_m < LOCAL_M; block_m++) {
int enable_input_d = block_k != 0;
TIME_EVENT(5);
barrier_wait(&barrier_smem_full[read_stage], read_phase);
TIME_EVENT(6);
unsigned long long smem_desc_sfa_base = make_smem_desc(&smem_sfa[read_stage][0], 16, 8 * 16, 0, 0, 0);
unsigned long long smem_desc_sfb_base = make_smem_desc(&smem_sfb[read_stage][0], 16, 8 * 16, 0, 0, 0);
unsigned long long smem_desc_a_base = make_smem_desc(&smem_a[read_stage][0], 1, 128 * 8, 0, 0, 2);
unsigned long long smem_desc_b_base = make_smem_desc(&smem_b[read_stage][0], 1, 128 * 8, read_stage, 0, 2);
unsigned int smem_desc_sfa_hi = make_warp_uniform((unsigned int) (smem_desc_sfa_base >> 32));
unsigned int smem_desc_sfa_lo = make_warp_uniform((unsigned int) (smem_desc_sfa_base >> 0));
unsigned int smem_desc_sfb_hi = make_warp_uniform((unsigned int) (smem_desc_sfb_base >> 32));
unsigned int smem_desc_sfb_lo = make_warp_uniform((unsigned int) (smem_desc_sfb_base >> 0));
unsigned int smem_desc_a_hi = make_warp_uniform((unsigned int) (smem_desc_a_base >> 32));
unsigned int smem_desc_a_lo = make_warp_uniform((unsigned int) (smem_desc_a_base >> 0));
unsigned int smem_desc_b_hi = make_warp_uniform((unsigned int) (smem_desc_b_base >> 32));
unsigned int smem_desc_b_lo = make_warp_uniform((unsigned int) (smem_desc_b_base >> 0));
// unsigned int insn_desc = (1u << 7) | (1u << 10) | (1u << 17) | (1u << 27); // e2m1, e2m1, N=8, M=128
int tmem_d = make_warp_uniform(block_m * 16);
int tmem_sfa_base = make_warp_uniform(32 + read_stage * 8 * (BLOCK_K / 2 / 2 / sizeof(uint4)));
static_assert(BLOCK_K == 256);
ASMV({
.reg .pred elect;
.reg .pred pred;
.reg .b32 tmem_d, tmem_sfa_base, enable_input_d, barrier;
.reg .b32 sfa_lo_start, sfb_lo_start, a_lo_start, b_lo_start;
.reg .b32 sfa_hi, sfb_hi, a_hi, b_hi;
.reg .b32 sfa_lo, sfb_lo, a_lo, b_lo;
.reg .b32 idesc;
.reg .b64 sfa, sfb, a, b;
mov.b32 tmem_d, %0;
mov.b32 tmem_sfa_base, %1;
mov.b32 enable_input_d, %2;
mov.b32 barrier, %3;
mov.b32 sfa_lo_start, %4;
mov.b32 sfb_lo_start, %5;
mov.b32 a_lo_start, %6;
mov.b32 b_lo_start, %7;
mov.b32 sfa_hi, %8;
mov.b32 sfb_hi, %9;
mov.b32 a_hi, %10;
mov.b32 b_hi, %11;
elect.sync _|elect, -1;
add.u32 sfa_lo, sfa_lo_start, 0;
mov.b64 sfa, {sfa_lo, sfa_hi};
add.u32 sfb_lo, sfb_lo_start, 0;
mov.b64 sfb, {sfb_lo, sfb_hi};
@elect tcgen05.cp.cta_group::1.32x128b.warpx4 [tmem_sfa_base+ 0], sfa;
@elect tcgen05.cp.cta_group::1.32x128b.warpx4 [tmem_sfa_base+ 4], sfb;
add.u32 sfa_lo, sfa_lo_start, 32;
mov.b64 sfa, {sfa_lo, sfa_hi};
add.u32 sfb_lo, sfb_lo_start, 1;
mov.b64 sfb, {sfb_lo, sfb_hi};
@elect tcgen05.cp.cta_group::1.32x128b.warpx4 [tmem_sfa_base+ 8], sfa;
@elect tcgen05.cp.cta_group::1.32x128b.warpx4 [tmem_sfa_base+12], sfb;
add.u32 sfa_lo, sfa_lo_start, 64;
mov.b64 sfa, {sfa_lo, sfa_hi};
add.u32 sfb_lo, sfb_lo_start, 2;
mov.b64 sfb, {sfb_lo, sfb_hi};
@elect tcgen05.cp.cta_group::1.32x128b.warpx4 [tmem_sfa_base+16], sfa;
@elect tcgen05.cp.cta_group::1.32x128b.warpx4 [tmem_sfa_base+20], sfb;
add.u32 sfa_lo, sfa_lo_start, 96;
mov.b64 sfa, {sfa_lo, sfa_hi};
add.u32 sfb_lo, sfb_lo_start, 3;
mov.b64 sfb, {sfb_lo, sfb_hi};
@elect tcgen05.cp.cta_group::1.32x128b.warpx4 [tmem_sfa_base+24], sfa;
@elect tcgen05.cp.cta_group::1.32x128b.warpx4 [tmem_sfa_base+28], sfb;
setp.ne.b32 pred, enable_input_d, 0;
mov.b32 idesc, 0x8020480;
add.u32 a_lo, a_lo_start, 0;
mov.b64 a, {a_lo, a_hi};
add.u32 b_lo, b_lo_start, 0;
mov.b64 b, {b_lo, b_hi};
@elect tcgen05.mma.cta_group::1.kind::mxf4nvf4.block_scale.block16 [tmem_d], a, b, idesc, [tmem_sfa_base+ 0], [tmem_sfa_base+ 4], pred;
add.u32 a_lo, a_lo_start, 2;
mov.b64 a, {a_lo, a_hi};
add.u32 b_lo, b_lo_start, 2;
mov.b64 b, {b_lo, b_hi};
@elect tcgen05.mma.cta_group::1.kind::mxf4nvf4.block_scale.block16 [tmem_d], a, b, idesc, [tmem_sfa_base+ 8], [tmem_sfa_base+12], 1;
add.u32 a_lo, a_lo_start, 4;
mov.b64 a, {a_lo, a_hi};
add.u32 b_lo, b_lo_start, 4;
mov.b64 b, {b_lo, b_hi};
@elect tcgen05.mma.cta_group::1.kind::mxf4nvf4.block_scale.block16 [tmem_d], a, b, idesc, [tmem_sfa_base+16], [tmem_sfa_base+20], 1;
add.u32 a_lo, a_lo_start, 6;
mov.b64 a, {a_lo, a_hi};
add.u32 b_lo, b_lo_start, 6;
mov.b64 b, {b_lo, b_hi};
@elect tcgen05.mma.cta_group::1.kind::mxf4nvf4.block_scale.block16 [tmem_d], a, b, idesc, [tmem_sfa_base+24], [tmem_sfa_base+28], 1;
@elect tcgen05.commit.cta_group::1.mbarrier::arrive::one.b64 [%3];
}) CLO(:
"r"(tmem_d), "r"(tmem_sfa_base), "r"(enable_input_d), "r"((unsigned int) __cvta_generic_to_shared(&barrier_smem_empty[read_stage])),
"r"(smem_desc_sfa_lo), "r"(smem_desc_sfb_lo), "r"(smem_desc_a_lo), "r"(smem_desc_b_lo),
"r"(smem_desc_sfa_hi), "r"(smem_desc_sfb_hi), "r"(smem_desc_a_hi), "r"(smem_desc_b_hi)
);
// #pragma unroll
// for (int partition_k = 0; partition_k < BLOCK_K / 64; partition_k += 1) {
// int tmem_sfa = tmem_sfa_base + 8 * partition_k;
// int tmem_sfb = tmem_sfa + 4;
// unsigned long long smem_desc_sfa = smem_desc_sfa_base + partition_k * (128 * 4 / 16);
// unsigned long long smem_desc_sfb = smem_desc_sfb_base + partition_k * (16 / 16);
// ASMV({
// .reg .pred elect;
// elect.sync _|elect, 0xFFFFFFFF;
// @elect tcgen05.cp.cta_group::1.32x128b.warpx4 [%0], %1;
// @elect tcgen05.cp.cta_group::1.32x128b.warpx4 [%2], %3;
// }) CLO(: "r"(tmem_sfa), "l"(smem_desc_sfa), "r"(tmem_sfb), "l"(smem_desc_sfb));
// }
// #pragma unroll
// for (int partition_k = 0; partition_k < BLOCK_K / 64; partition_k += 1) {
// int tmem_sfa = tmem_sfa_base + 8 * partition_k;
// int tmem_sfb = tmem_sfa + 4;
// unsigned long long smem_desc_a = smem_desc_a_base + partition_k * (32 / 16);
// unsigned long long smem_desc_b = smem_desc_b_base + partition_k * (32 / 16);
// ASMV({
// .reg .pred pred;
// setp.ne.b32 pred, %6, 0;
// .reg .pred elect;
// elect.sync _|elect, 0xFFFFFFFF;
// @elect tcgen05.mma.cta_group::1.kind::mxf4nvf4.block_scale.block16 [%0], %1, %2, %3, [%4], [%5], pred;
// }) CLO(: "r"(tmem_d), "l"(smem_desc_a), "l"(smem_desc_b), "r"(insn_desc), "r"(tmem_sfa), "r"(tmem_sfb), "r"(enable_input_d));
// enable_input_d = 1;
// }
// ASMV({
// .reg .pred elect;
// elect.sync _|elect, 0xFFFFFFFF;
// @elect tcgen05.commit.cta_group::1.mbarrier::arrive::one.b64 [%0];
// }) CLO(: "r"((unsigned int) __cvta_generic_to_shared(&barrier_smem_empty[read_stage])));
read_stage += 1;
if (read_stage == STAGE_COUNT) {
read_stage = 0;
read_phase ^= 1;
}
}
}
ASMV({
.reg .pred elect;
elect.sync _|elect, 0xFFFFFFFF;
@elect tcgen05.commit.cta_group::1.mbarrier::arrive::one.b64 [%0];
}) CLO(: "r"((unsigned int) __cvta_generic_to_shared(&barrier)));
barrier_wait(&barrier, barrier_wait_phase);
} else if (is_reset && SPLIT_K > 1) {
for (int i = 2 * lane_id; i < BLOCK_M; i += 64) {
reinterpret_cast<int&>(ptr_c[i]) = 0;
}
__syncwarp();
if (lane_id == 0) {
int* sem = &semaphore[gridDim.x / SPLIT_K * blockIdx.y + blockIdx.x / SPLIT_K][0];
ASMV(red.release.gpu.global.inc.u32 [%0], %1;) CLO(: "l"(sem), "n"(SPLIT_K - 1) : "memory");
TIME_EVENT(7);
for (;;) {
int value = 1;
ASMV(ld.acquire.gpu.global.b32 %0, [%1];) CLO("=r"(value) : "l"(sem) : "memory");
if (value == 0) break;
}
TIME_EVENT(8);
}
__syncwarp();
} else if (is_tma) {
int load_stage = 0;
int load_phase = 2; // modify to not wait initially?
for (int block_k = 0; block_k < cutover_k; block_k += BLOCK_K) {
for (int block_m = 0; block_m < LOCAL_M; block_m++) {
load_stage += 1;
if (load_stage == STAGE_COUNT) {
load_stage = 0;
load_phase = load_phase ? 0 : 1;
}
}
}
ASMV(bar.sync 1, 160;) CLO();
for (int block_k = cutover_k; block_k < k; block_k += BLOCK_K) {
#pragma unroll
for (int block_m = 0; block_m < LOCAL_M; block_m++) {
if (load_phase != 2) {
barrier_wait(&barrier_smem_empty[load_stage], load_phase);
}
ASMV({
.reg .pred elect;
elect.sync _|elect, -1;
@elect cp.async.bulk.tensor.3d.shared::cta.global.tile.mbarrier::complete_tx::bytes [%1], [%0, {%4, %5, %6}], [%2];
@elect mbarrier.expect_tx.shared::cta.b64 [%2], %3;
}) CLO(: "l"(&desc_a), "r"((unsigned int) __cvta_generic_to_shared(&smem_a[load_stage][0])), "r"((unsigned int) __cvta_generic_to_shared(&barrier_smem_full[load_stage])), "n"(BLOCK_M * BLOCK_K / 2), "r"((block_offset_k * k + block_k) / 4 / 2), "r"(block_offset_m * BLOCK_M * LOCAL_M + block_m * BLOCK_M), "r"(blockIdx.y));
ASMV({
.reg .pred elect;
elect.sync _|elect, -1;
@elect cp.async.bulk.shared::cta.global.mbarrier::complete_tx::bytes [%1], [%0], %3, [%2];
@elect mbarrier.expect_tx.shared::cta.b64 [%2], %3;
}) CLO(: "l"(&ptr_b[block_k / 2]), "r"((unsigned int) __cvta_generic_to_shared(&smem_b[load_stage][0])), "r"((unsigned int) __cvta_generic_to_shared(&barrier_smem_full[load_stage])), "n"(BLOCK_K / 2));
ASMV({
.reg .pred elect;
elect.sync _|elect, -1;
@elect cp.async.bulk.shared::cta.global.mbarrier::complete_tx::bytes [%1], [%0], %3, [%2];
@elect mbarrier.arrive.expect_tx.shared::cta.b64 _, [%2], %3;
}) CLO(: "l"(&ptr_sfa[block_k / 64 * STRIDE_SFA_K + block_m * (BLOCK_M / 128) * STRIDE_SFA_M]), "r"((unsigned int) __cvta_generic_to_shared(&smem_sfa[load_stage][0])), "r"((unsigned int) __cvta_generic_to_shared(&barrier_smem_full[load_stage])), "n"(BLOCK_M * BLOCK_K / 16));
// if (has_sfb) {
// int partition_k = threadIdx.x * sizeof(sf_vec_type);
// // 128b loads
// }
if (lane_id < BLOCK_K / 16 / 4) {
copy_sf_vec(&smem_sfb[load_stage][4 * sizeof(sf_vec_type) * lane_id], &ptr_sfb[((block_k / 16 + lane_id * sizeof(sf_vec_type)) % 4) + (block_k / 64 + lane_id * sizeof(sf_vec_type) / (64 / 16)) * STRIDE_SFB_K]);
ASMV(cp.async.mbarrier.arrive.noinc.shared::cta.b64 [%0];) CLO(: "r"((unsigned int) __cvta_generic_to_shared(&barrier_smem_full[load_stage])));
}
ASMV({
.reg .pred elect;
elect.sync _|elect, -1;
@elect mbarrier.arrive.shared::cta.b64 _, [%0], %1;
}) CLO(: "r"((unsigned int) __cvta_generic_to_shared(&barrier_smem_full[load_stage])), "n"(128 - BLOCK_K / 16 / 4 - 1));
load_stage += 1;
if (load_stage == STAGE_COUNT) {
load_stage = 0;
load_phase = load_phase ? 0 : 1;
}
}
}
}
TIME_EVENT(9);
__syncthreads();
TIME_EVENT(10);
// asm volatile("griddepcontrol.launch_dependents;");
// load tmem
__attribute__((aligned(16))) __shared__ __half result[LOCAL_M * BLOCK_M];
if (is_load) {
#pragma unroll
for (int block_m = 0; block_m < LOCAL_M; block_m++) {
int tmem_d = block_m * 16;
float ld_result;
ASMV(tcgen05.ld.sync.aligned.32x32b.x1.b32 {%0}, [%1];) CLO("=f"(ld_result) : "r"(tmem_d));
result[block_m * BLOCK_M + threadIdx.x] = __float2half(ld_result);
}
}
TIME_EVENT(11);
__syncthreads();
TIME_EVENT(12);
if (is_mma) {
asm volatile("tcgen05.dealloc.cta_group::1.sync.aligned.b32 0, 512;");
}
if (is_load) {
for (int i = threadIdx.x * 8; i < LOCAL_M * BLOCK_M; i += 128 * 8) {
if (SPLIT_K > 1) {
ASM({
.reg .v4 .f16x2 reg;
ld.shared.v4.b32 reg, [%0];
red.relaxed.gpu.global.add.noftz.v4.f16x2 [%1], reg;
}) CLO(: "r"((unsigned int) __cvta_generic_to_shared(result + i)), "l"(ptr_c + i));
} else {
reinterpret_cast<uint4&>(ptr_c[i]) = reinterpret_cast<uint4&>(result[i]);
}
}
}
TIME_EVENT(13);
#ifdef TIMING
__syncthreads();
if (blockIdx.x == 0 && blockIdx.y == 0 && threadIdx.x == 0) {
for (int i = 0; i < 3; i++) {
for (int j = 0; j < 128; j += 1) {
printf("%d %d %d %ld\\n", i, j, tag[i][j], timing[i][j]);
}
}
}
#endif
}
void launch(
int m,
int l,
void * ptr_a,
void * ptr_sfa,
void * ptr_b,
void * ptr_sfb,
void * ptr_c
) {
dim3 grid(SPLIT_K * m / BLOCK_M / LOCAL_M, l);
dim3 block(256);
void* func = (void*) kernel;
cudaLaunchConfig_t launch_config = {0};
launch_config.blockDim = block;
launch_config.gridDim = grid;
// launch_config.dynamicSmemBytes = 0;
cudaLaunchAttribute attrs[16];
attrs[0].id = cudaLaunchAttributeIgnore;
#ifdef CLUSTER_M
attrs[0].id = cudaLaunchAttributeClusterDimension;
attrs[0].val.clusterDim.x = CLUSTER_M;
#endif
attrs[0].val.clusterDim.y = 1;
attrs[0].val.clusterDim.z = 1;
launch_config.attrs = attrs;
launch_config.numAttrs = 1;
CUtensorMap desc_a;
uint64_t global_dims[3] = {STATIC_K / 2 / 4, (uint64_t) m, (uint64_t) l};
uint64_t global_strides[2] = {STRIDE_A_M, STRIDE_A_L};
uint32_t box_dim[3] = {BLOCK_K / 2 / 4, BLOCK_M, 1};
uint32_t element_strides[3] = {1, 1, 1};
// todo try different L2 PROMOTION
cuTensorMapEncodeTiled(&desc_a, CU_TENSOR_MAP_DATA_TYPE_INT32, 3, ptr_a, global_dims, global_strides, box_dim, element_strides, CU_TENSOR_MAP_INTERLEAVE_NONE, CU_TENSOR_MAP_SWIZZLE_128B, CU_TENSOR_MAP_L2_PROMOTION_NONE, CU_TENSOR_MAP_FLOAT_OOB_FILL_NONE);
void* args[] = {
(void*) &ptr_a,
(void*) &ptr_sfa,
(void*) &ptr_b,
(void*) &ptr_sfb,
(void*) &ptr_c,
(void*) &desc_a
};
cudaLaunchKernelExC(&launch_config, func, args);
}
#ifdef BUILD_PYTORCH
void run(torch::Tensor a, torch::Tensor b, torch::Tensor sfa, torch::Tensor sfb, torch::Tensor c) {
int m = a.size(0);
int k = 2 * a.size(1);
int l = a.size(2);
TORCH_CHECK(k == STATIC_K, "K must match static shape");
TORCH_CHECK(m % (BLOCK_M * LOCAL_M) == 0, "m must divide blocking evenly");
TORCH_CHECK(k % (BLOCK_K * SPLIT_K) == 0, "k must divide blocking evenly");
TORCH_CHECK(a.dim() == 3, "A must be 3d");
TORCH_CHECK(a.stride(0) == STRIDE_A_M, "M stride A must match static shape");
TORCH_CHECK(a.stride(1) == 1, "K stride in A must be 1");
TORCH_CHECK(a.stride(2) == STRIDE_A_L, "L stride A must match static shape");
TORCH_CHECK(b.dim() == 3, "B must be 3d");
TORCH_CHECK(b.stride(1) == 1, "K stride in B must be 1");
TORCH_CHECK(b.stride(2) == STRIDE_B_L, "L stride B must match static shape");
TORCH_CHECK(sfa.dim() == 6, "SFA must be 6d");
TORCH_CHECK(sfa.stride(0) == 16, "M0 stride in SFA must be 16");
TORCH_CHECK(sfa.stride(1) == 4, "M1 stride in SFA must be 4");
TORCH_CHECK(sfa.stride(2) == STRIDE_SFA_M, "M2 stride in SFA must match static shape");
TORCH_CHECK(sfa.stride(3) == 1, "K0 stride in SFA must be 1");
TORCH_CHECK(sfa.stride(4) == STRIDE_SFA_K, "K1 stride in SFA must match static shape");
TORCH_CHECK(sfa.stride(5) == STRIDE_SFA_L, "L stride in SFA must match static shape");
TORCH_CHECK(sfb.dim() == 6, "SFB must be 6d");
TORCH_CHECK(sfb.stride(0) == 16, "M0 stride in SFB must be 16");
TORCH_CHECK(sfb.stride(1) == 4, "M1 stride in SFB must be 4");
TORCH_CHECK(sfb.stride(3) == 1, "K0 stride in SFA must be 1");
TORCH_CHECK(sfb.stride(4) == STRIDE_SFB_K, "K1 stride in SFA must match static shape");
TORCH_CHECK(sfb.stride(5) == STRIDE_SFB_L, "L stride in SFA must match static shape");
TORCH_CHECK(c.dim() == 3, "C must be 3d");
TORCH_CHECK(c.stride(0) == 1, "M stride in C must be 1");
TORCH_CHECK(c.stride(2) == STRIDE_C_L, "L stride in C must match static shape");
launch(
m,
l,
a.data_ptr(),
sfa.data_ptr(),
b.data_ptr(),
sfb.data_ptr(),
c.data_ptr()
);
}
#endif
"""
],
cpp_sources=[
"void run(torch::Tensor, torch::Tensor, torch::Tensor, torch::Tensor, torch::Tensor);"
],
name="inline_module_" + str(abs(hash(key))),
functions=["run"],
extra_cflags=["-DBUILD_PYTORCH", "-O3"] + defines,
extra_cuda_cflags=[
"-DBUILD_PYTORCH",
"--resource-usage",
"-gencode=arch=compute_100a,code=sm_100a",
"-O3",
]
+ defines,
extra_ldflags=["-LcublasLt", "-lcuda"],
)
PREPOPULATE = {
(
((7168, 8192, 1), (8192, 1, 58720256)),
((128, 8192, 1), (8192, 1, 1048576)),
((32, 4, 56, 4, 256, 1), (16, 4, 131072, 1, 512, 7340032)),
((32, 4, 1, 4, 256, 1), (16, 4, 131072, 1, 512, 131072)),
((7168, 1, 1), (1, 1, 7168)),
): {"LOCAL_M": 1, "USE_PDL": 1, "CLUSTER_M": None, "SPLIT_K": 2, "STAGE_COUNT": 6},
(
((4096, 3584, 8), (3584, 1, 14680064)),
((128, 3584, 8), (3584, 1, 458752)),
((32, 4, 32, 4, 112, 8), (16, 4, 57344, 1, 512, 1835008)),
((32, 4, 1, 4, 112, 8), (16, 4, 57344, 1, 512, 57344)),
((4096, 1, 8), (1, 1, 4096)),
): {"LOCAL_M": 2, "USE_PDL": 1, "CLUSTER_M": None, "SPLIT_K": 1, "STAGE_COUNT": 6},
(
((7168, 1024, 4), (1024, 1, 7340032)),
((128, 1024, 4), (1024, 1, 131072)),
((32, 4, 56, 4, 32, 4), (16, 4, 16384, 1, 512, 917504)),
((32, 4, 1, 4, 32, 4), (16, 4, 16384, 1, 512, 16384)),
((7168, 1, 4), (1, 1, 7168)),
): {"LOCAL_M": 2, "USE_PDL": 1, "CLUSTER_M": None, "SPLIT_K": 1, "STAGE_COUNT": 6},
}
kernels = {}
for k, tunable in PREPOPULATE.items():
kernels[k] = compile_kernel(k, tunable)
kernels[k].run(*[torch.empty_strided(size, stride, dtype=dtype, device='cuda') for (size, stride), dtype in zip(k, [torch.float4_e2m1fn_x2, torch.float4_e2m1fn_x2, torch.float8_e4m3fnuz, torch.float8_e4m3fnuz, torch.float16])])
torch.cuda.synchronize()
def custom_kernel(
data: input_t,
) -> output_t:
a, b, _, _, sfa_permuted, sfb_permuted, c = data
args = a, b, sfa_permuted, sfb_permuted, c
key = kernel_key(*args)
if key not in kernels:
kernels[key] = compile_kernel(
key, heuristic(a, b, sfa_permuted, sfb_permuted, c)
)
kernels[key].run(*args)
return c
scrolls · 758 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 97026.
⋯ 54 unchanged linesdef compile_kernel(key, tunables):- print(key, tunables)+ # print(key, tunables)((a_shape, a_stride),(_, b_stride),⋯ 19 unchanged linescuda_sources=["""#ifdef BUILD_PYTORCH- #include <c10/cuda/CUDAStream.h>#endif#include <cuda_fp4.h>#include <cuda_fp8.h>#include <cuda_fp16.h>#include <cuda_runtime.h>- #include <iostream>+ #include <cuda.h>#include <vector>+ #include <cstdint>#define ASM(x...) asm(#x#define ASMV(x...) asm volatile(#x⋯ 38 unchanged linesasm("cp.async.ca.shared.global [%0], [%1], 4;" :: "r"((unsigned int)__cvta_generic_to_shared(dst)), "l"(src) : "memory");};+ __device__ unsigned int make_warp_uniform(unsigned int v) {+ return __shfl_sync(0xffffffff, v, 0);+ }+__device__ void barrier_wait(unsigned long long* barrier, int barrier_wait_phase) {ASMV({.reg .pred done;⋯ 4 unchanged lines}) CLO(: "r"((unsigned int) __cvta_generic_to_shared(barrier)), "r"(barrier_wait_phase));}- __launch_bounds__(128 + 64, 1)+ #ifdef TIMING+ #define TIME_EVENT(event_id) do { \+ if (is_timer) { \+ timing[time_slot][current_slot] = clock(); \+ tag[time_slot][current_slot] = event_id; \+ current_slot += 1; \+ } \+ } while (0)+ #else+ #define TIME_EVENT(x) do {} while(0)+ #endif++ __launch_bounds__(256, 1)__global__ void kernel(const __nv_fp4x2_e2m1 * __restrict__ ptr_a,const __nv_fp8_e4m3 * __restrict__ ptr_sfa,const __nv_fp4x2_e2m1 * __restrict__ ptr_b,const __nv_fp8_e4m3 * __restrict__ ptr_sfb,- __half * __restrict__ ptr_c+ __half * __restrict__ ptr_c,+ __grid_constant__ const CUtensorMap desc_a) {+ #ifdef TIMING+ __shared__ clock_t timing[3][128];+ __shared__ int tag[3][128];+ bool is_timer = threadIdx.x == 0 || threadIdx.x == 128 || threadIdx.x == (128+32);+ int time_slot = threadIdx.x == 0 ? 0 : threadIdx.x == 128 ? 1 : 2;+ int current_slot = 0;+ TIME_EVENT(0);+ #endif++// let's do the "real" blocking: 128x64 in M and K, block along M and Lptr_a += blockIdx.y * STRIDE_A_L;ptr_b += blockIdx.y * STRIDE_B_L;⋯ 17 unchanged lines// have 128 threads, each computing one output element// SF: MN, K -> addr is ((32, 4, RM), (16, 4, RK)):((16, 4, ?), (0, 1, ?))- const int stride_k_rest = 2 * 9 * 16 * (BLOCK_M / 8 + 1);__attribute__((aligned(128))) __shared__ __nv_fp8_e4m3 smem_sfb[STAGE_COUNT][4 * BLOCK_K / 16];- __attribute__((aligned(128))) __shared__ __nv_fp4x2_e2m1 smem_b[STAGE_COUNT][BLOCK_K / 2];- // 9 * 16, 2 * 9 * 17- __attribute__((aligned(128))) __shared__ __nv_fp4x2_e2m1 smem_a[STAGE_COUNT][stride_k_rest * (BLOCK_K / 2 / 32)];+ __attribute__((aligned(1024))) __shared__ __nv_fp4x2_e2m1 smem_b[STAGE_COUNT][BLOCK_K / 2];__attribute__((aligned(128))) __shared__ __nv_fp8_e4m3 smem_sfa[STAGE_COUNT][BLOCK_M * (BLOCK_K / 16)];+ // 9 * 16, 2 * 9 * 17+ __attribute__((aligned(1024))) __shared__ __nv_fp4x2_e2m1 smem_a[STAGE_COUNT][BLOCK_M * (BLOCK_K / 2)];int warp_id = threadIdx.x / 32;int wg_id = warp_id / 4;bool is_load = wg_id == 0;bool is_mma = warp_id == 4;bool is_reset = warp_id == 5;+ bool is_tma = warp_id == 6;int lane_id = threadIdx.x % 32;__shared__ int tmem_addr;- if (is_mma) {- asm volatile("tcgen05.alloc.cta_group::1.sync.aligned.shared::cta.b32 [%0], 512;" :: "r"((unsigned int) __cvta_generic_to_shared(&tmem_addr)));- }__shared__ unsigned long long barrier;if (warp_id == 0 && lane_id == 0) {asm volatile("mbarrier.init.shared::cta.b64 [%0], 1;" :: "r"((unsigned int) __cvta_generic_to_shared(&barrier)));⋯ 7 unchanged linesif (warp_id == 2 && lane_id < STAGE_COUNT) {asm volatile("mbarrier.init.shared::cta.b64 [%0], 128;" :: "r"((unsigned int) __cvta_generic_to_shared(&barrier_smem_full[lane_id])));}+ if (warp_id == 3) {+ ASMV(prefetch.tensormap [%0];) CLO(: "l"(&desc_a));+ }+ TIME_EVENT(1);__syncthreads();+ TIME_EVENT(2);auto mde = [](unsigned int x) { return (unsigned long long) ((x & 0x3FFFF) >> 4); };auto make_smem_desc = [mde](void* smem_ptr, unsigned long long leading_dimension, unsigned long long stride_dimension, unsigned long long matrix_base_offset, unsigned long long is_leading_byte_address_absolute, unsigned long long swizzling_mode) {return mde(__cvta_generic_to_shared(smem_ptr)) | (mde(leading_dimension) << 16) | (mde(stride_dimension) << 32) | (1ull << 46) | (matrix_base_offset << 49) | (is_leading_byte_address_absolute << 52) | (swizzling_mode << 61);};- const int vec_size_elems = 32;- const int threads_in_k = BLOCK_K / vec_size_elems;- const int threads_in_m = BLOCK_M / threads_in_k;- using vec_type = uint4;- int ptr_a_m_offset = threadIdx.x / threads_in_k;- int ptr_a_k_offset = (threadIdx.x % threads_in_k) * sizeof(vec_type);- ptr_a += ptr_a_k_offset + ptr_a_m_offset * STRIDE_A_M;- int ptr_b_n_offset = threadIdx.x / threads_in_k;- int ptr_b_k_offset = (threadIdx.x % threads_in_k) * sizeof(vec_type);- ptr_b += ptr_b_k_offset;-- using sf_vec_type = unsigned int;-auto offset_a = [&](int m, int k) {- return k % 16 + (m % 8) * 16 + ((k / 16) % 2) * 9 * 16 + (m / 8) * 2 * 9 * 16 + (k / 32) * stride_k_rest;+ return (k % 16) + 16 * (((k / 16) % 8) ^ (m % 8)) + 128 * m;};+ using sf_vec_type = unsigned int;+ int cutover_k = (STAGE_COUNT / LOCAL_M) * BLOCK_K;+ if (cutover_k > k) cutover_k = k;+ cutover_k = 0;+if (is_load) {++ const int vec_size_elems = 32;+ const int threads_in_k = BLOCK_K / vec_size_elems;+ const int threads_in_m = BLOCK_M / threads_in_k;+ using vec_type = uint4;+ int ptr_a_m_offset = threadIdx.x / threads_in_k;+ int ptr_a_k_offset = (threadIdx.x % threads_in_k) * sizeof(vec_type);+ ptr_a += ptr_a_k_offset + ptr_a_m_offset * STRIDE_A_M;++ int ptr_b_n_offset = threadIdx.x / threads_in_k;+ int ptr_b_k_offset = (threadIdx.x % threads_in_k) * sizeof(vec_type);++ auto smem_a_offs = reinterpret_cast<__nv_fp4x2_e2m1 (*)[BLOCK_M*BLOCK_K/2]>(&smem_a[0][offset_a(ptr_a_m_offset, ptr_a_k_offset)]);+ // ptr_sfa += threadIdx.x * 16;+ // auto smem_sfa_offs = reinterpret_cast<__nv_fp8_e4m3 (*)[BLOCK_M*BLOCK_K/16]>(&smem_a[0][(threadIdx.x % 32) * 16 + 128 * 4 * (threadIdx.x / 32)]);++ int has_sfb = threadIdx.x * sizeof(sf_vec_type) < BLOCK_K / 16;++ ptr_b += ptr_b_k_offset;+int load_stage = 0;int load_phase = 2; // modify to not wait initially?asm volatile("griddepcontrol.wait;" ::: "memory");- #pragma unroll 1- for (int block_k = 0; block_k < k; block_k += BLOCK_K) {+ for (int block_k = 0; block_k < cutover_k; block_k += BLOCK_K) {#pragma unrollfor (int block_m = 0; block_m < LOCAL_M; block_m++) {+ TIME_EVENT(3);if (load_phase != 2) {barrier_wait(&barrier_smem_empty[load_stage], load_phase);}+ TIME_EVENT(4);#pragma unrollfor (int partition_m = 0; partition_m < BLOCK_M; partition_m += threads_in_m) {- copy_b_vec(&smem_a[load_stage][offset_a(partition_m + ptr_a_m_offset, ptr_a_k_offset)], &ptr_a[block_k / 2 + partition_m * STRIDE_A_M + block_m * BLOCK_M * STRIDE_A_M]);+ copy_a_vec(&smem_a_offs[load_stage][partition_m * 128], &ptr_a[block_k / 2 + partition_m * STRIDE_A_M + block_m * BLOCK_M * STRIDE_A_M]);}- // load bif (ptr_b_n_offset == 0) {copy_b_vec(&smem_b[load_stage][ptr_b_k_offset], &ptr_b[block_k / 2]);}- // load sf a- // new storage format is 4K x 32M4 x 4M1 x RK- // 4 * 32 * 4 / 16 = 32, RK = BLOCK_K / 16 / 4#pragma unrollfor (int rest_k = threadIdx.x / 32; rest_k < BLOCK_K / 16 / 4; rest_k += BLOCK_M / 32) {copy_a_vec(&smem_sfa[load_stage][(threadIdx.x % 32) * 16 + 128 * 4 * rest_k], &ptr_sfa[(threadIdx.x % 32) * 16 + (rest_k + block_k / 64) * STRIDE_SFA_K + block_m * (BLOCK_M / 128) * STRIDE_SFA_M]);}- // load sf b- #pragma unroll- for (int partition_k = threadIdx.x * sizeof(sf_vec_type); partition_k < BLOCK_K / 16; partition_k += BLOCK_M * sizeof(sf_vec_type)) {- // 128b loads+ if (has_sfb) {+ int partition_k = threadIdx.x * sizeof(sf_vec_type);copy_sf_vec(&smem_sfb[load_stage][4 * partition_k], &ptr_sfb[((block_k / 16 + partition_k) % 4) + (block_k / 64 + partition_k / (64 / 16)) * STRIDE_SFB_K]);}ASMV(cp.async.mbarrier.arrive.noinc.shared::cta.b64 [%0];) CLO(: "r"((unsigned int) __cvta_generic_to_shared(&barrier_smem_full[load_stage])));⋯ 4 unchanged lines}}}+ ASMV(bar.arrive 1, 160;) CLO();+ // start computing} else if (is_mma) {+ asm volatile("tcgen05.alloc.cta_group::1.sync.aligned.shared::cta.b32 [%0], 512;" :: "r"((unsigned int) __cvta_generic_to_shared(&tmem_addr)));int read_stage = 0;int read_phase = 0;- #pragma unroll 1for (int block_k = 0; block_k < k; block_k += BLOCK_K) {#pragma unrollfor (int block_m = 0; block_m < LOCAL_M; block_m++) {int enable_input_d = block_k != 0;++ TIME_EVENT(5);barrier_wait(&barrier_smem_full[read_stage], read_phase);+ TIME_EVENT(6);unsigned long long smem_desc_sfa_base = make_smem_desc(&smem_sfa[read_stage][0], 16, 8 * 16, 0, 0, 0);unsigned long long smem_desc_sfb_base = make_smem_desc(&smem_sfb[read_stage][0], 16, 8 * 16, 0, 0, 0);- unsigned long long smem_desc_a_base = make_smem_desc(&smem_a[read_stage][0], 16 * 9, 16 * 9 * 2, 0, 0, 0);- unsigned long long smem_desc_b_base = make_smem_desc(&smem_b[read_stage][0], 16, 0, 0, 0, 0);- unsigned int insn_desc = (1u << 7) | (1u << 10) | (1u << 17) | (1u << 27); // e2m1, e2m1, N=8, M=128- #pragma unroll- for (int partition_k = 0; partition_k < BLOCK_K / 2; partition_k += 2 * sizeof(uint4)) {- int tmem_sfa = 32 + 8 * (partition_k / 2 / sizeof(uint4)) + read_stage * 8 * (BLOCK_K / 2 / 2 / sizeof(uint4));- int tmem_sfb = tmem_sfa + 4;- unsigned long long smem_desc_sfa = smem_desc_sfa_base + mde((partition_k / 32) * 128 * 4);- unsigned long long smem_desc_sfb = smem_desc_sfb_base + mde((partition_k / 32) * 16);- ASMV({- .reg .pred elect;- elect.sync _|elect, 0xFFFFFFFF;- @elect tcgen05.cp.cta_group::1.32x128b.warpx4 [%0], %1;- @elect tcgen05.cp.cta_group::1.32x128b.warpx4 [%2], %3;- }) CLO(: "r"(tmem_sfa), "l"(smem_desc_sfa), "r"(tmem_sfb), "l"(smem_desc_sfb));- }- #pragma unroll- for (int partition_k = 0; partition_k < BLOCK_K / 2; partition_k += 2 * sizeof(uint4)) {- int tmem_d = block_m * 16;- int tmem_sfa = 32 + 8 * (partition_k / 2 / sizeof(uint4)) + read_stage * 8 * (BLOCK_K / 2 / 2 / sizeof(uint4));- int tmem_sfb = tmem_sfa + 4;- unsigned long long smem_desc_a = smem_desc_a_base + mde(offset_a(0, partition_k));- unsigned long long smem_desc_b = smem_desc_b_base + mde(partition_k);- ASMV({- .reg .pred pred;- setp.ne.b32 pred, %6, 0;- .reg .pred elect;- elect.sync _|elect, 0xFFFFFFFF;- @elect tcgen05.mma.cta_group::1.kind::mxf4nvf4.block_scale.block16 [%0], %1, %2, %3, [%4], [%5], pred;- }) CLO(: "r"(tmem_d), "l"(smem_desc_a), "l"(smem_desc_b), "r"(insn_desc), "r"(tmem_sfa), "r"(tmem_sfb), "r"(enable_input_d));- enable_input_d = 1;- }+ unsigned long long smem_desc_a_base = make_smem_desc(&smem_a[read_stage][0], 1, 128 * 8, 0, 0, 2);+ unsigned long long smem_desc_b_base = make_smem_desc(&smem_b[read_stage][0], 1, 128 * 8, read_stage, 0, 2);++ unsigned int smem_desc_sfa_hi = make_warp_uniform((unsigned int) (smem_desc_sfa_base >> 32));+ unsigned int smem_desc_sfa_lo = make_warp_uniform((unsigned int) (smem_desc_sfa_base >> 0));+ unsigned int smem_desc_sfb_hi = make_warp_uniform((unsigned int) (smem_desc_sfb_base >> 32));+ unsigned int smem_desc_sfb_lo = make_warp_uniform((unsigned int) (smem_desc_sfb_base >> 0));++ unsigned int smem_desc_a_hi = make_warp_uniform((unsigned int) (smem_desc_a_base >> 32));+ unsigned int smem_desc_a_lo = make_warp_uniform((unsigned int) (smem_desc_a_base >> 0));+ unsigned int smem_desc_b_hi = make_warp_uniform((unsigned int) (smem_desc_b_base >> 32));+ unsigned int smem_desc_b_lo = make_warp_uniform((unsigned int) (smem_desc_b_base >> 0));++ // unsigned int insn_desc = (1u << 7) | (1u << 10) | (1u << 17) | (1u << 27); // e2m1, e2m1, N=8, M=128+ int tmem_d = make_warp_uniform(block_m * 16);+ int tmem_sfa_base = make_warp_uniform(32 + read_stage * 8 * (BLOCK_K / 2 / 2 / sizeof(uint4)));++ static_assert(BLOCK_K == 256);+ASMV({.reg .pred elect;- elect.sync _|elect, 0xFFFFFFFF;- @elect tcgen05.commit.cta_group::1.mbarrier::arrive::one.b64 [%0];- }) CLO(: "r"((unsigned int) __cvta_generic_to_shared(&barrier_smem_empty[read_stage])));+ .reg .pred pred;+ .reg .b32 tmem_d, tmem_sfa_base, enable_input_d, barrier;+ .reg .b32 sfa_lo_start, sfb_lo_start, a_lo_start, b_lo_start;+ .reg .b32 sfa_hi, sfb_hi, a_hi, b_hi;+ .reg .b32 sfa_lo, sfb_lo, a_lo, b_lo;+ .reg .b32 idesc;+ .reg .b64 sfa, sfb, a, b;+ mov.b32 tmem_d, %0;+ mov.b32 tmem_sfa_base, %1;+ mov.b32 enable_input_d, %2;+ mov.b32 barrier, %3;+ mov.b32 sfa_lo_start, %4;+ mov.b32 sfb_lo_start, %5;+ mov.b32 a_lo_start, %6;+ mov.b32 b_lo_start, %7;+ mov.b32 sfa_hi, %8;+ mov.b32 sfb_hi, %9;+ mov.b32 a_hi, %10;+ mov.b32 b_hi, %11;+ elect.sync _|elect, -1;++ add.u32 sfa_lo, sfa_lo_start, 0;+ mov.b64 sfa, {sfa_lo, sfa_hi};+ add.u32 sfb_lo, sfb_lo_start, 0;+ mov.b64 sfb, {sfb_lo, sfb_hi};+ @elect tcgen05.cp.cta_group::1.32x128b.warpx4 [tmem_sfa_base+ 0], sfa;+ @elect tcgen05.cp.cta_group::1.32x128b.warpx4 [tmem_sfa_base+ 4], sfb;++ add.u32 sfa_lo, sfa_lo_start, 32;+ mov.b64 sfa, {sfa_lo, sfa_hi};+ add.u32 sfb_lo, sfb_lo_start, 1;+ mov.b64 sfb, {sfb_lo, sfb_hi};+ @elect tcgen05.cp.cta_group::1.32x128b.warpx4 [tmem_sfa_base+ 8], sfa;+ @elect tcgen05.cp.cta_group::1.32x128b.warpx4 [tmem_sfa_base+12], sfb;++ add.u32 sfa_lo, sfa_lo_start, 64;+ mov.b64 sfa, {sfa_lo, sfa_hi};+ add.u32 sfb_lo, sfb_lo_start, 2;+ mov.b64 sfb, {sfb_lo, sfb_hi};+ @elect tcgen05.cp.cta_group::1.32x128b.warpx4 [tmem_sfa_base+16], sfa;+ @elect tcgen05.cp.cta_group::1.32x128b.warpx4 [tmem_sfa_base+20], sfb;++ add.u32 sfa_lo, sfa_lo_start, 96;+ mov.b64 sfa, {sfa_lo, sfa_hi};+ add.u32 sfb_lo, sfb_lo_start, 3;+ mov.b64 sfb, {sfb_lo, sfb_hi};+ @elect tcgen05.cp.cta_group::1.32x128b.warpx4 [tmem_sfa_base+24], sfa;+ @elect tcgen05.cp.cta_group::1.32x128b.warpx4 [tmem_sfa_base+28], sfb;++ setp.ne.b32 pred, enable_input_d, 0;+ mov.b32 idesc, 0x8020480;++ add.u32 a_lo, a_lo_start, 0;+ mov.b64 a, {a_lo, a_hi};+ add.u32 b_lo, b_lo_start, 0;+ mov.b64 b, {b_lo, b_hi};+ @elect tcgen05.mma.cta_group::1.kind::mxf4nvf4.block_scale.block16 [tmem_d], a, b, idesc, [tmem_sfa_base+ 0], [tmem_sfa_base+ 4], pred;++ add.u32 a_lo, a_lo_start, 2;+ mov.b64 a, {a_lo, a_hi};+ add.u32 b_lo, b_lo_start, 2;+ mov.b64 b, {b_lo, b_hi};+ @elect tcgen05.mma.cta_group::1.kind::mxf4nvf4.block_scale.block16 [tmem_d], a, b, idesc, [tmem_sfa_base+ 8], [tmem_sfa_base+12], 1;++ add.u32 a_lo, a_lo_start, 4;+ mov.b64 a, {a_lo, a_hi};+ add.u32 b_lo, b_lo_start, 4;+ mov.b64 b, {b_lo, b_hi};+ @elect tcgen05.mma.cta_group::1.kind::mxf4nvf4.block_scale.block16 [tmem_d], a, b, idesc, [tmem_sfa_base+16], [tmem_sfa_base+20], 1;++ add.u32 a_lo, a_lo_start, 6;+ mov.b64 a, {a_lo, a_hi};+ add.u32 b_lo, b_lo_start, 6;+ mov.b64 b, {b_lo, b_hi};+ @elect tcgen05.mma.cta_group::1.kind::mxf4nvf4.block_scale.block16 [tmem_d], a, b, idesc, [tmem_sfa_base+24], [tmem_sfa_base+28], 1;++ @elect tcgen05.commit.cta_group::1.mbarrier::arrive::one.b64 [%3];+ }) CLO(:+ "r"(tmem_d), "r"(tmem_sfa_base), "r"(enable_input_d), "r"((unsigned int) __cvta_generic_to_shared(&barrier_smem_empty[read_stage])),+ "r"(smem_desc_sfa_lo), "r"(smem_desc_sfb_lo), "r"(smem_desc_a_lo), "r"(smem_desc_b_lo),+ "r"(smem_desc_sfa_hi), "r"(smem_desc_sfb_hi), "r"(smem_desc_a_hi), "r"(smem_desc_b_hi)+ );+ // #pragma unroll+ // for (int partition_k = 0; partition_k < BLOCK_K / 64; partition_k += 1) {+ // int tmem_sfa = tmem_sfa_base + 8 * partition_k;+ // int tmem_sfb = tmem_sfa + 4;+ // unsigned long long smem_desc_sfa = smem_desc_sfa_base + partition_k * (128 * 4 / 16);+ // unsigned long long smem_desc_sfb = smem_desc_sfb_base + partition_k * (16 / 16);+ // ASMV({+ // .reg .pred elect;+ // elect.sync _|elect, 0xFFFFFFFF;+ // @elect tcgen05.cp.cta_group::1.32x128b.warpx4 [%0], %1;+ // @elect tcgen05.cp.cta_group::1.32x128b.warpx4 [%2], %3;+ // }) CLO(: "r"(tmem_sfa), "l"(smem_desc_sfa), "r"(tmem_sfb), "l"(smem_desc_sfb));+ // }+ // #pragma unroll+ // for (int partition_k = 0; partition_k < BLOCK_K / 64; partition_k += 1) {+ // int tmem_sfa = tmem_sfa_base + 8 * partition_k;+ // int tmem_sfb = tmem_sfa + 4;+ // unsigned long long smem_desc_a = smem_desc_a_base + partition_k * (32 / 16);+ // unsigned long long smem_desc_b = smem_desc_b_base + partition_k * (32 / 16);+ // ASMV({+ // .reg .pred pred;+ // setp.ne.b32 pred, %6, 0;+ // .reg .pred elect;+ // elect.sync _|elect, 0xFFFFFFFF;+ // @elect tcgen05.mma.cta_group::1.kind::mxf4nvf4.block_scale.block16 [%0], %1, %2, %3, [%4], [%5], pred;+ // }) CLO(: "r"(tmem_d), "l"(smem_desc_a), "l"(smem_desc_b), "r"(insn_desc), "r"(tmem_sfa), "r"(tmem_sfb), "r"(enable_input_d));+ // enable_input_d = 1;+ // }+ // ASMV({+ // .reg .pred elect;+ // elect.sync _|elect, 0xFFFFFFFF;+ // @elect tcgen05.commit.cta_group::1.mbarrier::arrive::one.b64 [%0];+ // }) CLO(: "r"((unsigned int) __cvta_generic_to_shared(&barrier_smem_empty[read_stage])));read_stage += 1;if (read_stage == STAGE_COUNT) {⋯ 16 unchanged linesif (lane_id == 0) {int* sem = &semaphore[gridDim.x / SPLIT_K * blockIdx.y + blockIdx.x / SPLIT_K][0];ASMV(red.release.gpu.global.inc.u32 [%0], %1;) CLO(: "l"(sem), "n"(SPLIT_K - 1) : "memory");+ TIME_EVENT(7);for (;;) {int value = 1;ASMV(ld.acquire.gpu.global.b32 %0, [%1];) CLO("=r"(value) : "l"(sem) : "memory");if (value == 0) break;}+ TIME_EVENT(8);}__syncwarp();+ } else if (is_tma) {+ int load_stage = 0;+ int load_phase = 2; // modify to not wait initially?+ for (int block_k = 0; block_k < cutover_k; block_k += BLOCK_K) {+ for (int block_m = 0; block_m < LOCAL_M; block_m++) {+ load_stage += 1;+ if (load_stage == STAGE_COUNT) {+ load_stage = 0;+ load_phase = load_phase ? 0 : 1;+ }+ }+ }+ ASMV(bar.sync 1, 160;) CLO();+ for (int block_k = cutover_k; block_k < k; block_k += BLOCK_K) {+ #pragma unroll+ for (int block_m = 0; block_m < LOCAL_M; block_m++) {+ if (load_phase != 2) {+ barrier_wait(&barrier_smem_empty[load_stage], load_phase);+ }+ ASMV({+ .reg .pred elect;+ elect.sync _|elect, -1;+ @elect cp.async.bulk.tensor.3d.shared::cta.global.tile.mbarrier::complete_tx::bytes [%1], [%0, {%4, %5, %6}], [%2];+ @elect mbarrier.expect_tx.shared::cta.b64 [%2], %3;+ }) CLO(: "l"(&desc_a), "r"((unsigned int) __cvta_generic_to_shared(&smem_a[load_stage][0])), "r"((unsigned int) __cvta_generic_to_shared(&barrier_smem_full[load_stage])), "n"(BLOCK_M * BLOCK_K / 2), "r"((block_offset_k * k + block_k) / 4 / 2), "r"(block_offset_m * BLOCK_M * LOCAL_M + block_m * BLOCK_M), "r"(blockIdx.y));+ ASMV({+ .reg .pred elect;+ elect.sync _|elect, -1;+ @elect cp.async.bulk.shared::cta.global.mbarrier::complete_tx::bytes [%1], [%0], %3, [%2];+ @elect mbarrier.expect_tx.shared::cta.b64 [%2], %3;+ }) CLO(: "l"(&ptr_b[block_k / 2]), "r"((unsigned int) __cvta_generic_to_shared(&smem_b[load_stage][0])), "r"((unsigned int) __cvta_generic_to_shared(&barrier_smem_full[load_stage])), "n"(BLOCK_K / 2));++ ASMV({+ .reg .pred elect;+ elect.sync _|elect, -1;+ @elect cp.async.bulk.shared::cta.global.mbarrier::complete_tx::bytes [%1], [%0], %3, [%2];+ @elect mbarrier.arrive.expect_tx.shared::cta.b64 _, [%2], %3;+ }) CLO(: "l"(&ptr_sfa[block_k / 64 * STRIDE_SFA_K + block_m * (BLOCK_M / 128) * STRIDE_SFA_M]), "r"((unsigned int) __cvta_generic_to_shared(&smem_sfa[load_stage][0])), "r"((unsigned int) __cvta_generic_to_shared(&barrier_smem_full[load_stage])), "n"(BLOCK_M * BLOCK_K / 16));+ // if (has_sfb) {+ // int partition_k = threadIdx.x * sizeof(sf_vec_type);+ // // 128b loads+ // }+ if (lane_id < BLOCK_K / 16 / 4) {+ copy_sf_vec(&smem_sfb[load_stage][4 * sizeof(sf_vec_type) * lane_id], &ptr_sfb[((block_k / 16 + lane_id * sizeof(sf_vec_type)) % 4) + (block_k / 64 + lane_id * sizeof(sf_vec_type) / (64 / 16)) * STRIDE_SFB_K]);+ ASMV(cp.async.mbarrier.arrive.noinc.shared::cta.b64 [%0];) CLO(: "r"((unsigned int) __cvta_generic_to_shared(&barrier_smem_full[load_stage])));+ }++ ASMV({+ .reg .pred elect;+ elect.sync _|elect, -1;+ @elect mbarrier.arrive.shared::cta.b64 _, [%0], %1;+ }) CLO(: "r"((unsigned int) __cvta_generic_to_shared(&barrier_smem_full[load_stage])), "n"(128 - BLOCK_K / 16 / 4 - 1));++ load_stage += 1;+ if (load_stage == STAGE_COUNT) {+ load_stage = 0;+ load_phase = load_phase ? 0 : 1;+ }+ }+ }}+ TIME_EVENT(9);__syncthreads();+ TIME_EVENT(10);// asm volatile("griddepcontrol.launch_dependents;");// load tmem⋯ 8 unchanged linesresult[block_m * BLOCK_M + threadIdx.x] = __float2half(ld_result);}}+ TIME_EVENT(11);__syncthreads();+ TIME_EVENT(12);if (is_mma) {asm volatile("tcgen05.dealloc.cta_group::1.sync.aligned.b32 0, 512;");}⋯ 10 unchanged lines}}}+ TIME_EVENT(13);+ #ifdef TIMING+ __syncthreads();+ if (blockIdx.x == 0 && blockIdx.y == 0 && threadIdx.x == 0) {+ for (int i = 0; i < 3; i++) {+ for (int j = 0; j < 128; j += 1) {+ printf("%d %d %d %ld\\n", i, j, tag[i][j], timing[i][j]);+ }+ }+ }+ #endif}void launch(⋯ 3 unchanged linesvoid * ptr_sfa,void * ptr_b,void * ptr_sfb,- void * ptr_c,- cudaStream_t stream+ void * ptr_c) {dim3 grid(SPLIT_K * m / BLOCK_M / LOCAL_M, l);- dim3 block(128 + 64);+ dim3 block(256);void* func = (void*) kernel;- cudaLaunchConfig_t launch_config;+ cudaLaunchConfig_t launch_config = {0};launch_config.blockDim = block;launch_config.gridDim = grid;- launch_config.stream = stream;- launch_config.dynamicSmemBytes = 0;+ // launch_config.dynamicSmemBytes = 0;cudaLaunchAttribute attrs[16];attrs[0].id = cudaLaunchAttributeIgnore;#ifdef CLUSTER_M⋯ 2 unchanged lines#endifattrs[0].val.clusterDim.y = 1;attrs[0].val.clusterDim.z = 1;- attrs[1].id = cudaLaunchAttributeProgrammaticStreamSerialization;- attrs[1].val.programmaticStreamSerializationAllowed = USE_PDL;launch_config.attrs = attrs;- launch_config.numAttrs = 2;+ launch_config.numAttrs = 1;+ CUtensorMap desc_a;+ uint64_t global_dims[3] = {STATIC_K / 2 / 4, (uint64_t) m, (uint64_t) l};+ uint64_t global_strides[2] = {STRIDE_A_M, STRIDE_A_L};+ uint32_t box_dim[3] = {BLOCK_K / 2 / 4, BLOCK_M, 1};+ uint32_t element_strides[3] = {1, 1, 1};+ // todo try different L2 PROMOTION+ cuTensorMapEncodeTiled(&desc_a, CU_TENSOR_MAP_DATA_TYPE_INT32, 3, ptr_a, global_dims, global_strides, box_dim, element_strides, CU_TENSOR_MAP_INTERLEAVE_NONE, CU_TENSOR_MAP_SWIZZLE_128B, CU_TENSOR_MAP_L2_PROMOTION_NONE, CU_TENSOR_MAP_FLOAT_OOB_FILL_NONE);void* args[] = {(void*) &ptr_a,(void*) &ptr_sfa,(void*) &ptr_b,(void*) &ptr_sfb,(void*) &ptr_c,+ (void*) &desc_a};cudaLaunchKernelExC(&launch_config, func, args);}⋯ 42 unchanged linessfa.data_ptr(),b.data_ptr(),sfb.data_ptr(),- c.data_ptr(),- c10::cuda::getCurrentCUDAStream().stream()+ c.data_ptr());}#endif⋯ 12 unchanged lines"-O3",]+ defines,- extra_ldflags=["-LcublasLt"],+ extra_ldflags=["-LcublasLt", "-lcuda"],)⋯ 24 unchanged linesfor k, tunable in PREPOPULATE.items():kernels[k] = compile_kernel(k, tunable)kernels[k].run(*[torch.empty_strided(size, stride, dtype=dtype, device='cuda') for (size, stride), dtype in zip(k, [torch.float4_e2m1fn_x2, torch.float4_e2m1fn_x2, torch.float8_e4m3fnuz, torch.float8_e4m3fnuz, torch.float16])])+ torch.cuda.synchronize()def custom_kernel(
scrolls · 586 diff lines total
Best evidence level for this revision: reported
JSON