submission 111027
v0i0 · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 1002 lines, June 9 Researcher Reciprocity License v1.0.
submission-28-12-20-32.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-nvfp4-gemv-111027?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:638359569e141628b5b85b7001f44d19e8d2995dbe2ae221e4e45a8346d3d1f0
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() {fp4
uint4 veca = *reinterpret_cast<uint4*>(&smem_a[read_stage][offset_a(local_m, partition_k)]); // this has 32 fp4 elemsfp8
__half_raw half_value = __nv_cvt_fp8_to_halfraw(reinterpret_cast<const __nv_fp8_storage_t&>(value), __NV_E4M3);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-28-12-20-32.py1002 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));
}
__device__ float e2m1x2_to_float(__nv_fp4x2_e2m1 value, int subbyte_idx) {
__half2_raw values = __nv_cvt_fp4x2_to_halfraw2(reinterpret_cast<const __nv_fp4x2_storage_t&>(value), __NV_E2M1);
return __half2float(reinterpret_cast<const __half&>(subbyte_idx == 0 ? values.x : values.y));
}
__device__ __half2 e2m1x2_to_half2(__nv_fp4x2_e2m1 value) {
__half2_raw values = __nv_cvt_fp4x2_to_halfraw2(reinterpret_cast<const __nv_fp4x2_storage_t&>(value), __NV_E2M1);
return reinterpret_cast<const __half2&>(values);
}
__device__ __half2 e2m1x2_to_half2(unsigned char value) {
__half2_raw values = __nv_cvt_fp4x2_to_halfraw2(reinterpret_cast<const __nv_fp4x2_storage_t&>(value), __NV_E2M1);
return reinterpret_cast<const __half2&>(values);
}
__device__ float e4m3_to_float(unsigned char value) {
__half_raw half_value = __nv_cvt_fp8_to_halfraw(reinterpret_cast<const __nv_fp8_storage_t&>(value), __NV_E4M3);
return __half2float(reinterpret_cast<const __half&>(half_value));
}
__device__ float e4m3_to_float(__nv_fp8_e4m3 value) {
__half_raw half_value = __nv_cvt_fp8_to_halfraw(reinterpret_cast<const __nv_fp8_storage_t&>(value), __NV_E4M3);
return __half2float(reinterpret_cast<const __half&>(half_value));
}
__device__ __half e4m3_to_half(__nv_fp8_e4m3 value) {
__half_raw half_value = __nv_cvt_fp8_to_halfraw(reinterpret_cast<const __nv_fp8_storage_t&>(value), __NV_E4M3);
return reinterpret_cast<const __half&>(half_value);
}
__device__ __half e4m3_to_half(unsigned char value) {
return e4m3_to_half(reinterpret_cast<const __nv_fp8_e4m3&>(value));
}
#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
struct TileIterator {
int start_m, start_k;
int num_tiles;
int total_k_tiles;
int total_m_tiles;
int total_sms;
int index;
// simple, we have N tiles
// we need to be able to decode them
// locally, we need to know this tile and the next tile
// we need to be able to chunk
//
};
__device__ int get_chunked_local_pos(int num, int idx, int total) {
int count = total / num;
int rest = total % num;
int pos = min(idx, rest) + idx * count;
return pos;
}
__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)];
__attribute__((aligned(16))) __shared__ __half result[LOCAL_M * BLOCK_M];
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();
#if 0
// start computing
int read_stage = 0;
int read_phase = 0;
int local_m = threadIdx.x;
float acc[LOCAL_M][8] = {0};
int idx = 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++) {
idx += 1;
if (idx % 2 == 0) {
read_stage += 1;
if (read_stage == STAGE_COUNT) {
read_stage = 0;
read_phase ^= 1;
}
continue;
}
barrier_wait(&barrier_smem_full[read_stage], read_phase);
// each thread reads one column of K values from A and B
// that means 256 / 32 = 8 elems, or 4 bytes.
// seems tunable, too.
// if i wanted 16 bytes, i'd have 8 threads along K
// so four along M, so B is only reused 8 times. that might be fine.
// load B: 32 elems, dequantize to fp16
// load sfB: 2 elem, dequantize to fp16
// load A: 32 elems in K, 8 elems in M, dequantize to fp16
// load sfB: 2 elem, dequantize to fp16
// update each accumulator
// do cutover: take every other one to tc or to local
int partition_k = (lane_id % 8) * 16;
// each thread loads 16 bytes == 32 elements
// for blocking 256, there are 8 (= 256 / 32) threads in the k dimension then
// this means there are 4 threads in the m dimension
// this means each thread then has 32 / 4 = 8 elements in the m dimension
// let them be contiguous
uint4 vecb = *reinterpret_cast<uint4*>(&smem_b[read_stage][partition_k]);
__half2 vecbv[sizeof(vecb)];
#pragma unroll
for (int i = 0; i < sizeof(vecb); i++) {
vecbv[i] = e2m1x2_to_half2(reinterpret_cast<unsigned char*>(&vecb)[i]);
}
uchar2 vecsfb = *reinterpret_cast<uchar2*>(&smem_sfb[read_stage][16 * (partition_k / 32) + (partition_k / 8) % 4]);
__half value_sfb[2] = {e4m3_to_half(vecsfb.x), e4m3_to_half(vecsfb.y)};
#pragma unroll
for (int the_m = 0; the_m < 8; the_m += 1) {
int local_m = warp_id * 32 + the_m + (lane_id / 8) * 8;
uint4 veca = *reinterpret_cast<uint4*>(&smem_a[read_stage][offset_a(local_m, partition_k)]); // this has 32 fp4 elems
__half2 vecav[sizeof(veca)];
#pragma unroll
for (int i = 0; i < sizeof(veca); i++) {
vecav[i] = e2m1x2_to_half2(reinterpret_cast<unsigned char*>(&veca)[i]);
}
uchar2 vecsfa = *reinterpret_cast<uchar2*>(&smem_sfa[read_stage][(partition_k / 8 / 4) * 128 * 4 + ((partition_k / 8) % 4) + (local_m % 32) * 16 + (local_m / 32) * 4]);
__half value_sfa[2] = {e4m3_to_half(vecsfa.x), e4m3_to_half(vecsfa.y)};
__half2 local_acc[2];
local_acc[0] = __hmul2(vecav[0], vecbv[0]);
local_acc[1] = __hmul2(vecav[8], vecbv[8]);
#pragma unroll
for (int i = 1; i < 8; i++) {
local_acc[0] = __hfma2(vecav[i+0], vecbv[i+0], local_acc[0]);
local_acc[1] = __hfma2(vecav[i+8], vecbv[i+8], local_acc[1]);
}
acc[block_m][the_m] += __half2float(__hmul(__hadd(local_acc[0].x, local_acc[0].y), __hmul(value_sfa[0], value_sfb[0])));
acc[block_m][the_m] += __half2float(__hmul(__hadd(local_acc[1].x, local_acc[1].y), __hmul(value_sfa[1], value_sfb[1])));
}
#if 0
// better maybe: each warp still has 32 elements, but each thread in the warp does, too, and they are split along K
#pragma unroll
for (int partition_k = 0; partition_k < BLOCK_K / 2; partition_k += 2 * sizeof(uint4)) {
uint4 veca = *reinterpret_cast<uint4*>(&smem_a[read_stage][offset_a(local_m, partition_k)]);
uint4 vecb = *reinterpret_cast<uint4*>(&smem_b[read_stage][partition_k]);
uchar4 vecsfa = *reinterpret_cast<uchar4*>(&smem_sfa[read_stage][(partition_k / 8 / 4) * 128 * 4 + (local_m % 32) * 16 + (local_m / 32) * 4]);
uchar4 vecsfb = *reinterpret_cast<uchar4*>(&smem_sfb[read_stage][4 * partition_k / 8]);
uint4 vecc = *reinterpret_cast<uint4*>(&smem_a[read_stage][offset_a(local_m, partition_k + sizeof(uint4))]);
uint4 vecd = *reinterpret_cast<uint4*>(&smem_b[read_stage][partition_k + sizeof(uint4)]);
auto dot_sf = [](unsigned char sfa, unsigned char sfb, unsigned int a0, unsigned int a1, unsigned int b0, unsigned int b1) {
uchar4 veca = reinterpret_cast<const uchar4&>(a0);
uchar4 vecb = reinterpret_cast<const uchar4&>(b0);
__half2 value_a = e2m1x2_to_half2(veca.x);
__half2 value_b = e2m1x2_to_half2(vecb.x);
__half2 acc = __hmul2(value_a, value_b);
value_a = e2m1x2_to_half2(veca.y);
value_b = e2m1x2_to_half2(vecb.y);
acc = __hfma2(value_a, value_b, acc);
value_a = e2m1x2_to_half2(veca.z);
value_b = e2m1x2_to_half2(vecb.z);
acc = __hfma2(value_a, value_b, acc);
value_a = e2m1x2_to_half2(veca.w);
value_b = e2m1x2_to_half2(vecb.w);
acc = __hfma2(value_a, value_b, acc);
veca = reinterpret_cast<const uchar4&>(a1);
vecb = reinterpret_cast<const uchar4&>(b1);
value_a = e2m1x2_to_half2(veca.x);
value_b = e2m1x2_to_half2(vecb.x);
acc = __hfma2(value_a, value_b, acc);
value_a = e2m1x2_to_half2(veca.y);
value_b = e2m1x2_to_half2(vecb.y);
acc = __hfma2(value_a, value_b, acc);
value_a = e2m1x2_to_half2(veca.z);
value_b = e2m1x2_to_half2(vecb.z);
acc = __hfma2(value_a, value_b, acc);
value_a = e2m1x2_to_half2(veca.w);
value_b = e2m1x2_to_half2(vecb.w);
acc = __hfma2(value_a, value_b, acc);
__half sum = __hadd(acc.x, acc.y);
float value_sfa = e4m3_to_float(sfa);
float value_sfb = e4m3_to_float(sfb);
return __half2float(sum) * value_sfa * value_sfb;
};
acc[block_m] += dot_sf(vecsfa.x, vecsfb.x, veca.x, veca.y, vecb.x, vecb.y);
acc[block_m] += dot_sf(vecsfa.y, vecsfb.y, veca.z, veca.w, vecb.z, vecb.w);
acc[block_m] += dot_sf(vecsfa.z, vecsfb.z, vecc.x, vecc.y, vecd.x, vecd.y);
acc[block_m] += dot_sf(vecsfa.w, vecsfb.w, vecc.z, vecc.w, vecd.z, vecd.w);
}
#endif
ASMV(bar.sync 2, 128;) CLO();
if (threadIdx.x == 0) {
ASMV(mbarrier.arrive.shared::cta.b64 _, [%0], 1;) 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;
}
}
}
for (int block_m = 0; block_m < LOCAL_M; block_m++) {
for (int i = 0; i < 8; i++) {
// reduce the accumulator
float reduced_acc = acc[block_m][i];
reduced_acc += __shfl_xor_sync((unsigned int) -1, reduced_acc, 1);
reduced_acc += __shfl_xor_sync((unsigned int) -1, reduced_acc, 2);
reduced_acc += __shfl_xor_sync((unsigned int) -1, reduced_acc, 4);
if (threadIdx.x % 8 == 0) {
result[block_m * BLOCK_M + 32 * warp_id + i + 8 * (lane_id / 8)] = __float2half(reduced_acc);
}
}
}
// __syncthreads();
// 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]);
// }
// }
#endif
} 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;
int idx = 0;
int has_zerod[LOCAL_M] = {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++) {
// idx += 1;
// if (idx % 2 == 1) {
// read_stage += 1;
// if (read_stage == STAGE_COUNT) {
// read_stage = 0;
// read_phase ^= 1;
// }
// continue;
// }
int enable_input_d = has_zerod[block_m];
has_zerod[block_m] = 1;
// 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
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] = __hadd(result[block_m * BLOCK_M + threadIdx.x], __float2half(ld_result));
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 · 1002 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 109661.
⋯ 147 unchanged lines}) CLO(: "r"((unsigned int) __cvta_generic_to_shared(barrier)), "r"(barrier_wait_phase));}+ __device__ float e2m1x2_to_float(__nv_fp4x2_e2m1 value, int subbyte_idx) {+ __half2_raw values = __nv_cvt_fp4x2_to_halfraw2(reinterpret_cast<const __nv_fp4x2_storage_t&>(value), __NV_E2M1);+ return __half2float(reinterpret_cast<const __half&>(subbyte_idx == 0 ? values.x : values.y));+ }++ __device__ __half2 e2m1x2_to_half2(__nv_fp4x2_e2m1 value) {+ __half2_raw values = __nv_cvt_fp4x2_to_halfraw2(reinterpret_cast<const __nv_fp4x2_storage_t&>(value), __NV_E2M1);+ return reinterpret_cast<const __half2&>(values);+ }++ __device__ __half2 e2m1x2_to_half2(unsigned char value) {+ __half2_raw values = __nv_cvt_fp4x2_to_halfraw2(reinterpret_cast<const __nv_fp4x2_storage_t&>(value), __NV_E2M1);+ return reinterpret_cast<const __half2&>(values);+ }++ __device__ float e4m3_to_float(unsigned char value) {+ __half_raw half_value = __nv_cvt_fp8_to_halfraw(reinterpret_cast<const __nv_fp8_storage_t&>(value), __NV_E4M3);+ return __half2float(reinterpret_cast<const __half&>(half_value));+ }++ __device__ float e4m3_to_float(__nv_fp8_e4m3 value) {+ __half_raw half_value = __nv_cvt_fp8_to_halfraw(reinterpret_cast<const __nv_fp8_storage_t&>(value), __NV_E4M3);+ return __half2float(reinterpret_cast<const __half&>(half_value));+ }++ __device__ __half e4m3_to_half(__nv_fp8_e4m3 value) {+ __half_raw half_value = __nv_cvt_fp8_to_halfraw(reinterpret_cast<const __nv_fp8_storage_t&>(value), __NV_E4M3);+ return reinterpret_cast<const __half&>(half_value);+ }++ __device__ __half e4m3_to_half(unsigned char value) {+ return e4m3_to_half(reinterpret_cast<const __nv_fp8_e4m3&>(value));+ }+#ifdef TIMING#define TIME_EVENT(event_id) do { \if (is_timer) { \⋯ 6 unchanged lines#define TIME_EVENT(x) do {} while(0)#endif+ struct TileIterator {+ int start_m, start_k;+ int num_tiles;+ int total_k_tiles;+ int total_m_tiles;+ int total_sms;+ int index;+ // simple, we have N tiles+ // we need to be able to decode them+ // locally, we need to know this tile and the next tile+ // we need to be able to chunk+ //+ };++ __device__ int get_chunked_local_pos(int num, int idx, int total) {+ int count = total / num;+ int rest = total % num;+ int pos = min(idx, rest) + idx * count;+ return pos;+ }+__launch_bounds__(256, 1)__global__ void kernel(const __nv_fp4x2_e2m1 * __restrict__ ptr_a,⋯ 41 unchanged lines__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)];+ __attribute__((aligned(16))) __shared__ __half result[LOCAL_M * BLOCK_M];int warp_id = threadIdx.x / 32;int wg_id = warp_id / 4;⋯ 36 unchanged linesint cutover_k = (STAGE_COUNT / LOCAL_M) * BLOCK_K;if (cutover_k > k) cutover_k = k;- cutover_k = 0;+ // cutover_k = 0;if (is_load) {⋯ 51 unchanged lines}}ASMV(bar.arrive 1, 160;) CLO();+ #if 0// start computing+ int read_stage = 0;+ int read_phase = 0;+ int local_m = threadIdx.x;+ float acc[LOCAL_M][8] = {0};+ int idx = 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++) {++ idx += 1;+ if (idx % 2 == 0) {+ read_stage += 1;+ if (read_stage == STAGE_COUNT) {+ read_stage = 0;+ read_phase ^= 1;+ }+ continue;+ }++ barrier_wait(&barrier_smem_full[read_stage], read_phase);++ // each thread reads one column of K values from A and B+ // that means 256 / 32 = 8 elems, or 4 bytes.+ // seems tunable, too.+ // if i wanted 16 bytes, i'd have 8 threads along K+ // so four along M, so B is only reused 8 times. that might be fine.++ // load B: 32 elems, dequantize to fp16+ // load sfB: 2 elem, dequantize to fp16+ // load A: 32 elems in K, 8 elems in M, dequantize to fp16+ // load sfB: 2 elem, dequantize to fp16+ // update each accumulator++ // do cutover: take every other one to tc or to local+ int partition_k = (lane_id % 8) * 16;++ // each thread loads 16 bytes == 32 elements+ // for blocking 256, there are 8 (= 256 / 32) threads in the k dimension then+ // this means there are 4 threads in the m dimension+ // this means each thread then has 32 / 4 = 8 elements in the m dimension+ // let them be contiguous+ uint4 vecb = *reinterpret_cast<uint4*>(&smem_b[read_stage][partition_k]);+ __half2 vecbv[sizeof(vecb)];+ #pragma unroll+ for (int i = 0; i < sizeof(vecb); i++) {+ vecbv[i] = e2m1x2_to_half2(reinterpret_cast<unsigned char*>(&vecb)[i]);+ }++ uchar2 vecsfb = *reinterpret_cast<uchar2*>(&smem_sfb[read_stage][16 * (partition_k / 32) + (partition_k / 8) % 4]);++ __half value_sfb[2] = {e4m3_to_half(vecsfb.x), e4m3_to_half(vecsfb.y)};++ #pragma unroll+ for (int the_m = 0; the_m < 8; the_m += 1) {+ int local_m = warp_id * 32 + the_m + (lane_id / 8) * 8;+ uint4 veca = *reinterpret_cast<uint4*>(&smem_a[read_stage][offset_a(local_m, partition_k)]); // this has 32 fp4 elems+ __half2 vecav[sizeof(veca)];+ #pragma unroll+ for (int i = 0; i < sizeof(veca); i++) {+ vecav[i] = e2m1x2_to_half2(reinterpret_cast<unsigned char*>(&veca)[i]);+ }++ uchar2 vecsfa = *reinterpret_cast<uchar2*>(&smem_sfa[read_stage][(partition_k / 8 / 4) * 128 * 4 + ((partition_k / 8) % 4) + (local_m % 32) * 16 + (local_m / 32) * 4]);+ __half value_sfa[2] = {e4m3_to_half(vecsfa.x), e4m3_to_half(vecsfa.y)};++ __half2 local_acc[2];+ local_acc[0] = __hmul2(vecav[0], vecbv[0]);+ local_acc[1] = __hmul2(vecav[8], vecbv[8]);+ #pragma unroll+ for (int i = 1; i < 8; i++) {+ local_acc[0] = __hfma2(vecav[i+0], vecbv[i+0], local_acc[0]);+ local_acc[1] = __hfma2(vecav[i+8], vecbv[i+8], local_acc[1]);+ }+ acc[block_m][the_m] += __half2float(__hmul(__hadd(local_acc[0].x, local_acc[0].y), __hmul(value_sfa[0], value_sfb[0])));+ acc[block_m][the_m] += __half2float(__hmul(__hadd(local_acc[1].x, local_acc[1].y), __hmul(value_sfa[1], value_sfb[1])));+ }+++ #if 0++ // better maybe: each warp still has 32 elements, but each thread in the warp does, too, and they are split along K+ #pragma unroll+ for (int partition_k = 0; partition_k < BLOCK_K / 2; partition_k += 2 * sizeof(uint4)) {+ uint4 veca = *reinterpret_cast<uint4*>(&smem_a[read_stage][offset_a(local_m, partition_k)]);+ uint4 vecb = *reinterpret_cast<uint4*>(&smem_b[read_stage][partition_k]);+ uchar4 vecsfa = *reinterpret_cast<uchar4*>(&smem_sfa[read_stage][(partition_k / 8 / 4) * 128 * 4 + (local_m % 32) * 16 + (local_m / 32) * 4]);+ uchar4 vecsfb = *reinterpret_cast<uchar4*>(&smem_sfb[read_stage][4 * partition_k / 8]);++ uint4 vecc = *reinterpret_cast<uint4*>(&smem_a[read_stage][offset_a(local_m, partition_k + sizeof(uint4))]);+ uint4 vecd = *reinterpret_cast<uint4*>(&smem_b[read_stage][partition_k + sizeof(uint4)]);++ auto dot_sf = [](unsigned char sfa, unsigned char sfb, unsigned int a0, unsigned int a1, unsigned int b0, unsigned int b1) {+ uchar4 veca = reinterpret_cast<const uchar4&>(a0);+ uchar4 vecb = reinterpret_cast<const uchar4&>(b0);+ __half2 value_a = e2m1x2_to_half2(veca.x);+ __half2 value_b = e2m1x2_to_half2(vecb.x);+ __half2 acc = __hmul2(value_a, value_b);+ value_a = e2m1x2_to_half2(veca.y);+ value_b = e2m1x2_to_half2(vecb.y);+ acc = __hfma2(value_a, value_b, acc);+ value_a = e2m1x2_to_half2(veca.z);+ value_b = e2m1x2_to_half2(vecb.z);+ acc = __hfma2(value_a, value_b, acc);+ value_a = e2m1x2_to_half2(veca.w);+ value_b = e2m1x2_to_half2(vecb.w);+ acc = __hfma2(value_a, value_b, acc);+ veca = reinterpret_cast<const uchar4&>(a1);+ vecb = reinterpret_cast<const uchar4&>(b1);+ value_a = e2m1x2_to_half2(veca.x);+ value_b = e2m1x2_to_half2(vecb.x);+ acc = __hfma2(value_a, value_b, acc);+ value_a = e2m1x2_to_half2(veca.y);+ value_b = e2m1x2_to_half2(vecb.y);+ acc = __hfma2(value_a, value_b, acc);+ value_a = e2m1x2_to_half2(veca.z);+ value_b = e2m1x2_to_half2(vecb.z);+ acc = __hfma2(value_a, value_b, acc);+ value_a = e2m1x2_to_half2(veca.w);+ value_b = e2m1x2_to_half2(vecb.w);+ acc = __hfma2(value_a, value_b, acc);++ __half sum = __hadd(acc.x, acc.y);+ float value_sfa = e4m3_to_float(sfa);+ float value_sfb = e4m3_to_float(sfb);+ return __half2float(sum) * value_sfa * value_sfb;+ };++ acc[block_m] += dot_sf(vecsfa.x, vecsfb.x, veca.x, veca.y, vecb.x, vecb.y);+ acc[block_m] += dot_sf(vecsfa.y, vecsfb.y, veca.z, veca.w, vecb.z, vecb.w);+ acc[block_m] += dot_sf(vecsfa.z, vecsfb.z, vecc.x, vecc.y, vecd.x, vecd.y);+ acc[block_m] += dot_sf(vecsfa.w, vecsfb.w, vecc.z, vecc.w, vecd.z, vecd.w);+ }+ #endif++ ASMV(bar.sync 2, 128;) CLO();+ if (threadIdx.x == 0) {+ ASMV(mbarrier.arrive.shared::cta.b64 _, [%0], 1;) 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;+ }+ }+ }++ for (int block_m = 0; block_m < LOCAL_M; block_m++) {+ for (int i = 0; i < 8; i++) {+ // reduce the accumulator+ float reduced_acc = acc[block_m][i];+ reduced_acc += __shfl_xor_sync((unsigned int) -1, reduced_acc, 1);+ reduced_acc += __shfl_xor_sync((unsigned int) -1, reduced_acc, 2);+ reduced_acc += __shfl_xor_sync((unsigned int) -1, reduced_acc, 4);+ if (threadIdx.x % 8 == 0) {+ result[block_m * BLOCK_M + 32 * warp_id + i + 8 * (lane_id / 8)] = __float2half(reduced_acc);+ }+ }+ }+ // __syncthreads();+ // 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]);+ // }+ // }+ #endif} 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;+ int idx = 0;+ int has_zerod[LOCAL_M] = {0};for (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;+ // idx += 1;+ // if (idx % 2 == 1) {+ // read_stage += 1;+ // if (read_stage == STAGE_COUNT) {+ // read_stage = 0;+ // read_phase ^= 1;+ // }+ // continue;+ // }++ int enable_input_d = has_zerod[block_m];+ has_zerod[block_m] = 1;+ // int enable_input_d = block_k != 0;+TIME_EVENT(5);barrier_wait(&barrier_smem_full[read_stage], read_phase);TIME_EVENT(6);⋯ 235 unchanged lines// asm volatile("griddepcontrol.launch_dependents;");// load tmem- __attribute__((aligned(16))) __shared__ __half result[LOCAL_M * BLOCK_M];if (is_load) {#pragma unroll⋯ 1 unchanged linesint 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] = __hadd(result[block_m * BLOCK_M + threadIdx.x], __float2half(ld_result));result[block_m * BLOCK_M + threadIdx.x] = __float2half(ld_result);}}
scrolls · 308 diff lines total
Best evidence level for this revision: reported
JSON