submission 97026
v0i0 · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 537 lines, June 9 Researcher Reciprocity License v1.0.
submission-22-12-31-53.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-nvfp4-gemv-97026?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:f78afe6e0182e3cfc52d84c6a75e490e78f8fc4ed49ec3aa5be18e41b6e951c9
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
__attribute__((aligned(128))) __shared__ __nv_fp8_e4m3 smem_sfb[STAGE_COUNT][4 * BLOCK_K / 16];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;vector-width = uint4
using vec_type = uint4;Kernel source
submission-22-12-31-53.py537 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
#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 <vector>
#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__ 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));
}
__launch_bounds__(128 + 64, 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
) {
// 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, ?))
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(128))) __shared__ __nv_fp8_e4m3 smem_sfa[STAGE_COUNT][BLOCK_M * (BLOCK_K / 16)];
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;
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)));
}
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])));
}
__syncthreads();
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;
};
if (is_load) {
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) {
#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);
}
#pragma unroll
for (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]);
}
// load b
if (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 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]);
}
// 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
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;
}
}
}
} else if (is_mma) {
int read_stage = 0;
int read_phase = 0;
#pragma unroll 1
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;
barrier_wait(&barrier_smem_full[read_stage], read_phase);
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;
}
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");
for (;;) {
int value = 1;
ASMV(ld.acquire.gpu.global.b32 %0, [%1];) CLO("=r"(value) : "l"(sem) : "memory");
if (value == 0) break;
}
}
__syncwarp();
}
__syncthreads();
// 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);
}
}
__syncthreads();
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]);
}
}
}
}
void launch(
int m,
int l,
void * ptr_a,
void * ptr_sfa,
void * ptr_b,
void * ptr_sfb,
void * ptr_c,
cudaStream_t stream
) {
dim3 grid(SPLIT_K * m / BLOCK_M / LOCAL_M, l);
dim3 block(128 + 64);
void* func = (void*) kernel;
cudaLaunchConfig_t launch_config;
launch_config.blockDim = block;
launch_config.gridDim = grid;
launch_config.stream = stream;
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;
attrs[1].id = cudaLaunchAttributeProgrammaticStreamSerialization;
attrs[1].val.programmaticStreamSerializationAllowed = USE_PDL;
launch_config.attrs = attrs;
launch_config.numAttrs = 2;
void* args[] = {
(void*) &ptr_a,
(void*) &ptr_sfa,
(void*) &ptr_b,
(void*) &ptr_sfb,
(void*) &ptr_c,
};
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(),
c10::cuda::getCurrentCUDAStream().stream()
);
}
#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"],
)
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])])
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 · 537 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 88871.
import torchfrom task import input_t, output_tfrom torch.utils.cpp_extension import load_inline+ from triton.testing import do_bench+ import itertools- module = load_inline(cuda_sources=["""++ 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#include <c10/cuda/CUDAStream.h>#endif⋯ 9 unchanged lines#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 STAGE_COUNT = 8;const int BLOCK_K = 256;template<int N>⋯ 23 unchanged lines}) CLO(: "r"((unsigned int) __cvta_generic_to_shared(barrier)), "r"(barrier_wait_phase));}- template<int LOCAL_M>__launch_bounds__(128 + 64, 1)__global__ void kernel(- const int m,- const int k,- const int l,const __nv_fp4x2_e2m1 * __restrict__ ptr_a,- const int stride_a_m,- const int stride_a_l,const __nv_fp8_e4m3 * __restrict__ ptr_sfa,- const int stride_sfa_m,- const int stride_sfa_k,- const int stride_sfa_l,const __nv_fp4x2_e2m1 * __restrict__ ptr_b,- const int stride_b_l,const __nv_fp8_e4m3 * __restrict__ ptr_sfb,- const int stride_sfb_k,- const int stride_sfb_l,- __half * __restrict__ ptr_c,- const int stride_c_m,- const int stride_c_l,- int split_k+ __half * __restrict__ ptr_c) {// 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;+ 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;- ptr_a += blockIdx.x * BLOCK_M * stride_a_m * LOCAL_M;- ptr_sfa += blockIdx.x * (BLOCK_M / 128) * stride_sfa_m * LOCAL_M;- ptr_c += blockIdx.x * BLOCK_M * stride_c_m * LOCAL_M;+ int block_offset_m = blockIdx.x / SPLIT_K;+ int block_offset_k = blockIdx.x % SPLIT_K;+ int k = STATIC_K / SPLIT_K;- ptr_c += threadIdx.x * stride_c_m;+ 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, ?))⋯ 29 unchanged lines}__syncthreads();- auto make_smem_desc = [](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) {- auto mde = [](unsigned int x) { return (unsigned long long) ((x & 0x3FFFF) >> 4); };+ 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);};⋯ 3 unchanged linesusing 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;+ 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);⋯ 9 unchanged linesint load_stage = 0;int load_phase = 2; // modify to not wait initially?asm volatile("griddepcontrol.wait;" ::: "memory");+ #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++) {if (load_phase != 2) {barrier_wait(&barrier_smem_empty[load_stage], load_phase);}#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_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]);}// load bif (ptr_b_n_offset == 0) {⋯ 2 unchanged lines// 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]);+ 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 unrollfor (int partition_k = threadIdx.x * sizeof(sf_vec_type); partition_k < BLOCK_K / 16; partition_k += BLOCK_M * sizeof(sf_vec_type)) {// 128b loads- 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]);+ 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;⋯ 6 unchanged lines} else if (is_mma) {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;barrier_wait(&barrier_smem_full[read_stage], read_phase);++ 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 unrollfor (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;- barrier_wait(&barrier_smem_full[read_stage], read_phase);- unsigned long long smem_desc_sfa = make_smem_desc(&smem_sfa[read_stage][(partition_k / 32) * 128 * 4], 16, 8 * 16, 0, 0, 0);- unsigned long long smem_desc_sfb = make_smem_desc(&smem_sfb[read_stage][(partition_k / 32) * 16], 16, 8 * 16, 0, 0, 0);- unsigned long long smem_desc_a = make_smem_desc(&smem_a[read_stage][offset_a(0, partition_k)], 16 * 9, 16 * 9 * 2, 0, 0, 0);- unsigned long long smem_desc_b = make_smem_desc(&smem_b[read_stage][partition_k], 16, 0, 0, 0, 0);- unsigned int insn_desc = (1u << 7) | (1u << 10) | (1u << 17) | (1u << 27); // e2m1, e2m1, N=8, M=128+ 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;⋯ 22 unchanged lines@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");+ for (;;) {+ int value = 1;+ ASMV(ld.acquire.gpu.global.b32 %0, [%1];) CLO("=r"(value) : "l"(sem) : "memory");+ if (value == 0) break;+ }+ }+ __syncwarp();}__syncthreads();- asm volatile("griddepcontrol.launch_dependents;");+ // asm volatile("griddepcontrol.launch_dependents;");// load tmem- float result[LOCAL_M];+ __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] = ld_result;+ result[block_m * BLOCK_M + threadIdx.x] = __float2half(ld_result);}}__syncthreads();⋯ 1 unchanged linesasm volatile("tcgen05.dealloc.cta_group::1.sync.aligned.b32 0, 512;");}if (is_load) {- #pragma unroll- for (int block_m = 0; block_m < LOCAL_M; block_m++) {- ptr_c[block_m * BLOCK_M * stride_c_m] = __float2half(result[block_m]);+ 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]);+ }}}}void launch(int m,- int k,int l,void * ptr_a,- int stride_a_m,- int stride_a_l,void * ptr_sfa,- int stride_sfa_m,- int stride_sfa_k,- int stride_sfa_l,void * ptr_b,- int stride_b_l,void * ptr_sfb,- int stride_sfb_k,- int stride_sfb_l,void * ptr_c,- int stride_c_m,- int stride_c_l,cudaStream_t stream) {- int total_blocks = (m / BLOCK_M) * l;- int local_m = 1;- int split_k = 1;- if (total_blocks > 128 && (m % (BLOCK_M * 2)) == 0) {- local_m = 2;- }- // if (total_blocks < 64) {- // split_k = 2;- // }- dim3 grid(split_k * m / BLOCK_M / local_m, l);+ dim3 grid(SPLIT_K * m / BLOCK_M / LOCAL_M, l);dim3 block(128 + 64);- void* func = local_m == 2 ? (void*) kernel<2> : (void*) kernel<1>;+ void* func = (void*) kernel;cudaLaunchConfig_t launch_config;launch_config.blockDim = block;⋯ 1 unchanged lineslaunch_config.stream = stream;launch_config.dynamicSmemBytes = 0;cudaLaunchAttribute attrs[16];- attrs[0].id = cudaLaunchAttributeClusterDimension;attrs[0].id = cudaLaunchAttributeIgnore;- attrs[0].val.clusterDim.x = 1;+ #ifdef CLUSTER_M+ attrs[0].id = cudaLaunchAttributeClusterDimension;+ attrs[0].val.clusterDim.x = CLUSTER_M;+ #endifattrs[0].val.clusterDim.y = 1;attrs[0].val.clusterDim.z = 1;attrs[1].id = cudaLaunchAttributeProgrammaticStreamSerialization;- attrs[1].val.programmaticStreamSerializationAllowed = 1; // 1+ attrs[1].val.programmaticStreamSerializationAllowed = USE_PDL;launch_config.attrs = attrs;launch_config.numAttrs = 2;void* args[] = {- (void*) &m,- (void*) &k,- (void*) &l,(void*) &ptr_a,- (void*) &stride_a_m,- (void*) &stride_a_l,(void*) &ptr_sfa,- (void*) &stride_sfa_m,- (void*) &stride_sfa_k,- (void*) &stride_sfa_l,(void*) &ptr_b,- (void*) &stride_b_l,(void*) &ptr_sfb,- (void*) &stride_sfb_k,- (void*) &stride_sfb_l,(void*) &ptr_c,- (void*) &stride_c_m,- (void*) &stride_c_l,- (void*) &split_k};cudaLaunchKernelExC(&launch_config, func, args);}⋯ 3 unchanged linesint 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(b.stride(1) == 1, "K stride in B must be 1");- // 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(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,- k,l,a.data_ptr(),- a.stride(0),- a.stride(2),sfa.data_ptr(),- sfa.stride(2),- sfa.stride(4),- sfa.stride(5),b.data_ptr(),- b.stride(2),sfb.data_ptr(),- // sfb.stride(2),- sfb.stride(4),- sfb.stride(5),c.data_ptr(),- c.stride(0),- c.stride(2),c10::cuda::getCurrentCUDAStream().stream());}#endif- """],- cpp_sources=["void run(torch::Tensor, torch::Tensor, torch::Tensor, torch::Tensor, torch::Tensor);"],- name="inline_module",- functions=["run"],- extra_cflags=["-DBUILD_PYTORCH", "-O3"],- extra_cuda_cflags=["-DBUILD_PYTORCH", "--resource-usage", "-gencode=arch=compute_100a,code=sm_100a", "-O3"],- extra_ldflags=["-LcublasLt"],- )+ """+ ],+ 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"],+ )- # TODO- PRECOMPILE_AND_TUNE_STAGES = [- (16384, 1, 7168),- (7168, 8, 4096),- (2048, 4, 7168),- ]+ 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])])++def custom_kernel(data: input_t,) -> output_t:a, b, _, _, sfa_permuted, sfb_permuted, c = data- module.run(a, b, sfa_permuted, sfb_permuted, c)+ 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 · 554 diff lines total
Best evidence level for this revision: reported
JSON