submission 88871
v0i0 · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 403 lines, June 9 Researcher Reciprocity License v1.0.
submission-19-06-48-45.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-nvfp4-gemv-88871?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:b720e4fd2112f0fce84acdbae3f0e35eb906adb6bd654701ea7ac6c0887485fe
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
int split_ktcgen05
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-19-06-48-45.py403 lines
import torch
from task import input_t, output_t
from torch.utils.cpp_extension import load_inline
module = 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)
__device__ int semaphore[128][32] = {0};
const int BLOCK_M = 128;
const int STAGE_COUNT = 8;
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));
}
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
) {
// 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.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;
ptr_c += threadIdx.x * stride_c_m;
// 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 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); };
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");
for (int block_k = 0; block_k < k; block_k += BLOCK_K) {
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
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;
for (int block_k = 0; block_k < k; block_k += BLOCK_K) {
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);
#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;
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
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));
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);
}
__syncthreads();
asm volatile("griddepcontrol.launch_dependents;");
// load tmem
float result[LOCAL_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] = ld_result;
}
}
__syncthreads();
if (is_mma) {
asm 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]);
}
}
}
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 block(128 + 64);
void* func = local_m == 2 ? (void*) kernel<2> : (void*) kernel<1>;
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 = cudaLaunchAttributeClusterDimension;
attrs[0].id = cudaLaunchAttributeIgnore;
attrs[0].val.clusterDim.x = 1;
attrs[0].val.clusterDim.y = 1;
attrs[0].val.clusterDim.z = 1;
attrs[1].id = cudaLaunchAttributeProgrammaticStreamSerialization;
attrs[1].val.programmaticStreamSerializationAllowed = 1; // 1
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);
}
#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(a.dim() == 3, "A must be 3d");
TORCH_CHECK(a.stride(1) == 1, "K stride in A must be 1");
TORCH_CHECK(b.dim() == 3, "B must be 3d");
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(3) == 1, "K0 stride in SFA must be 1");
TORCH_CHECK(sfb.dim() == 6, "SFB must be 6d");
TORCH_CHECK(sfb.stride(3) == 1, "K0 stride in SFA must be 1");
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");
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"],
)
# TODO
PRECOMPILE_AND_TUNE_STAGES = [
(16384, 1, 7168),
(7168, 8, 4096),
(2048, 4, 7168),
]
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)
return c
scrolls · 403 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 74602.
⋯ 11 unchanged lines#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)++ __device__ int semaphore[128][32] = {0};+const int BLOCK_M = 128;const int STAGE_COUNT = 8;const int BLOCK_K = 256;- __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));+ template<int N>+ __device__ void cp_async_wait_group() {+ asm("cp.async.wait_group %0;" :: "n"(N) : "memory");}- __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__ 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__ __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__ 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__ 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__ 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__ 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__ 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__ __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);- }-- template<class V, class D, class S>- __device__ void copy(D* dst, const S* src) {- *reinterpret_cast<V*>(dst) = *reinterpret_cast<const V*>(src);- }-- template<int N>- __device__ void cp_async_wait_group() {- #define CAWG_COND(x) \- if constexpr (N == x) { \- asm("cp.async.wait_group " #x ";" ::: "memory"); \- }- CAWG_COND(16) CAWG_COND(15) CAWG_COND(14) CAWG_COND(13)- CAWG_COND(12) CAWG_COND(11) CAWG_COND(10) CAWG_COND( 9)- CAWG_COND( 8) CAWG_COND( 7) CAWG_COND( 6) CAWG_COND( 5)- CAWG_COND( 4) CAWG_COND( 3) CAWG_COND( 2) CAWG_COND( 1)- CAWG_COND( 0)- #undef CAWG_COND- }-- __launch_bounds__(BLOCK_M, 1)+ template<int LOCAL_M>+ __launch_bounds__(128 + 64, 1)__global__ void kernel(const int m,const int k,⋯ 12 unchanged linesconst int stride_sfb_l,__half * __restrict__ ptr_c,const int stride_c_m,- const int stride_c_l+ const int stride_c_l,+ int split_k) {// let's do the "real" blocking: 128x64 in M and K, block along M and Lptr_a += blockIdx.y * stride_a_l;⋯ 2 unchanged linesptr_sfb += blockIdx.y * stride_sfb_l;ptr_c += blockIdx.y * stride_c_l;- ptr_a += blockIdx.x * BLOCK_M * stride_a_m;- ptr_sfa += blockIdx.x * (BLOCK_M / 128) * stride_sfa_m;- ptr_c += blockIdx.x * BLOCK_M * stride_c_m;+ 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;+ ptr_c += threadIdx.x * stride_c_m;+// have 128 threads, each computing one output element- int local_m = threadIdx.x;// SF: MN, K -> addr is ((32, 4, RM), (16, 4, RK)):((16, 4, ?), (0, 1, ?))- ptr_c += local_m * stride_c_m;- __attribute__((aligned(16))) __shared__ __nv_fp4x2_e2m1 smem_a[STAGE_COUNT][BLOCK_M][BLOCK_K / 2 + 16];- __attribute__((aligned(16))) __shared__ __nv_fp8_e4m3 smem_sfa[STAGE_COUNT][BLOCK_M][BLOCK_K / 16];- __attribute__((aligned(16))) __shared__ __nv_fp4x2_e2m1 smem_b[STAGE_COUNT][BLOCK_K / 2 + 16];- __attribute__((aligned(16))) __shared__ __nv_fp8_e4m3 smem_sfb[STAGE_COUNT][BLOCK_K / 16];+ 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 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); };+ 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;⋯ 8 unchanged linesusing sf_vec_type = unsigned int;- float acc_out = 0;- int load_stage = 0;- int load_block_k = 0;- int read_stage = 0;- auto load_step = [&]() {- if (load_block_k < k) {- auto copy_a_vec = [](auto dst, auto src) {- asm("cp.async.cg.shared.global [%0], [%1], 16;" :: "r"((unsigned int)__cvta_generic_to_shared(dst)), "l"(src) : "memory");- };+ 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;+ };- auto copy_b_vec = [](auto dst, auto src) {- asm("cp.async.ca.shared.global [%0], [%1], 16;" :: "r"((unsigned int)__cvta_generic_to_shared(dst)), "l"(src) : "memory");- };-- auto copy_sf_vec = [](auto dst, auto src) {- asm("cp.async.ca.shared.global [%0], [%1], 4;" :: "r"((unsigned int)__cvta_generic_to_shared(dst)), "l"(src) : "memory");- };-- // load a- #pragma unroll- for (int partition_m = 0; partition_m < BLOCK_M; partition_m += threads_in_m) {- copy_a_vec(&smem_a[load_stage][partition_m + ptr_a_m_offset][ptr_a_k_offset], &ptr_a[load_block_k / 2 + partition_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[load_block_k / 2]);- }- // load sf a- for (int partition_m = threadIdx.x; partition_m < BLOCK_M; partition_m += BLOCK_M) {+ if (is_load) {+ 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 < k; block_k += BLOCK_K) {+ 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_k = 0; partition_k < BLOCK_K / 16; partition_k += sizeof(sf_vec_type)) {+ 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+ 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- // 256 elems = 16 sfs, 4 contigous, then 32, then 4, 4 sfs = 4B- copy_sf_vec(&smem_sfa[load_stage][partition_m][partition_k], &ptr_sfa[((load_block_k / 16 + partition_k) % 4) + (load_block_k / 64 + partition_k / (64 / 16)) * stride_sfa_k + (local_m % 32) * 16 + ((local_m / 32) % 4) * 4 + (local_m / 128) * stride_sfa_m]);+ 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;+ }}- // load sf b- // 256 elems = 16 scale factors = 1 ldgsts- #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][partition_k], &ptr_sfb[((load_block_k / 16 + partition_k) % 4) + (load_block_k / 64 + partition_k / (64 / 16)) * stride_sfb_k]);+ }+ } else if (is_mma) {+ int read_stage = 0;+ int read_phase = 0;+ for (int block_k = 0; block_k < k; block_k += BLOCK_K) {+ 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);+ #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;+ 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+ 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));+ 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;+ }}- load_block_k += BLOCK_K;- load_stage += 1;- load_stage %= STAGE_COUNT;}- asm("cp.async.commit_group;" ::: "memory");- };-- asm volatile("griddepcontrol.wait;" ::: "memory");- for (int i = 0; i < STAGE_COUNT - 1; i++) {- load_step();+ 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);}- for (int block_k = 0; block_k < k; block_k += BLOCK_K) {- // if constexpr (STAGE_COUNT == 4) {- // asm("cp.async.wait_group 2;" ::: "memory");- // }- cp_async_wait_group<STAGE_COUNT-2>();- __syncthreads();- load_step();- // __syncthreads();- #pragma unroll- for (int partition_k = 0; partition_k < BLOCK_K / 2; partition_k += sizeof(uint4)) {- // load at 64b from smem- uint4 veca = *reinterpret_cast<uint4*>(&smem_a[read_stage][local_m][partition_k]);- uint4 vecb = *reinterpret_cast<uint4*>(&smem_b[read_stage][partition_k]);- uchar2 vecsfa = *reinterpret_cast<uchar2*>(&smem_sfa[read_stage][local_m][partition_k / 8]);- uchar2 vecsfb = *reinterpret_cast<uchar2*>(&smem_sfb[read_stage][partition_k / 8]);+ __syncthreads();- 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);+ asm volatile("griddepcontrol.launch_dependents;");+ // load tmem+ float result[LOCAL_M];+ if (is_load) {- __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_out += dot_sf(vecsfa.x, vecsfb.x, veca.x, veca.y, vecb.x, vecb.y);- acc_out += dot_sf(vecsfa.y, vecsfb.y, veca.z, veca.w, vecb.z, vecb.w);+ #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] = ld_result;}-- read_stage += 1;- read_stage %= STAGE_COUNT;}- asm volatile("griddepcontrol.launch_dependents;" ::: "memory");- *ptr_c = __float2half(acc_out);+ __syncthreads();+ if (is_mma) {+ asm 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]);+ }+ }}void launch(⋯ 17 unchanged linesint stride_c_l,cudaStream_t stream) {- dim3 grid(m / BLOCK_M, l);- dim3 block(BLOCK_M);- void* func = (void*) kernel;- int dyn_smem = 0; //200*1024;- // cudaFuncSetAttribute(func, cudaFuncAttributeNonPortableClusterSizeAllowed, 1);- // if (dyn_smem > (48 * 1024)) {- // cudaFuncSetAttribute(func, cudaFuncAttributeMaxDynamicSharedMemorySize, dyn_smem);+ 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 block(128 + 64);+ void* func = local_m == 2 ? (void*) kernel<2> : (void*) kernel<1>;cudaLaunchConfig_t launch_config;launch_config.blockDim = block;launch_config.gridDim = grid;launch_config.stream = stream;- launch_config.dynamicSmemBytes = dyn_smem;+ launch_config.dynamicSmemBytes = 0;cudaLaunchAttribute attrs[16];attrs[0].id = cudaLaunchAttributeClusterDimension;attrs[0].id = cudaLaunchAttributeIgnore;⋯ 22 unchanged lines(void*) &stride_sfb_l,(void*) &ptr_c,(void*) &stride_c_m,- (void*) &stride_c_l+ (void*) &stride_c_l,+ (void*) &split_k};cudaLaunchKernelExC(&launch_config, func, args);}⋯ 49 unchanged linesextra_ldflags=["-LcublasLt"],)+ # TODO+ PRECOMPILE_AND_TUNE_STAGES = [+ (16384, 1, 7168),+ (7168, 8, 4096),+ (2048, 4, 7168),+ ]+def custom_kernel(data: input_t,) -> output_t:
scrolls · 450 diff lines total
Best evidence level for this revision: reported
JSON