submission 78065
gau.nernst · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 342 lines, June 9 Researcher Reciprocity License v1.0.
submission_v2_cpasync.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-nvfp4-gemv-78065?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:2eacadee22c6ae6fdff26a4240b5836b4974bae33d6a39d1141a14823aaa1e25
license declaredunknown
license concludedunknown
authorsgau.nernst
imported2026-08-15
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
async-copy
void cp_async_2d(int dst, const T *src, int src_stride, int tid) {fp4
"cvt.rn.f16x2.e2m1x2 %0, tmp0; // PTX only supports FP4->FP16\n"fp8
const __nv_fp8_e4m3 *SFA_ptr, // [L, M, K/8]num-warps = 4
constexpr int NUM_WARPS = 4;shared-memory
extern __shared__ char smem[];stages = 1
template <int THREAD_K, int NUM_STAGES = 1>vector-width = float2
float2 SFA_fp32x2[THREAD_M], SFB_fp32x2;Kernel source
submission_v2_cpasync.py342 lines
#!POPCORN leaderboard nvfp4_gemv
from pathlib import Path
import torch
from task import input_t, output_t
from torch.utils.cpp_extension import load_inline
CUDA_SRC = r"""
#include <cuda_fp16.h>
#include <cuda_fp8.h>
#include <torch/library.h>
#include <ATen/ATen.h>
#include <ATen/core/Tensor.h>
#include <ATen/cuda/CUDAUtils.h>
#include <ATen/cuda/CUDAContext.h>
constexpr int WARP_SIZE = 32;
constexpr int NUM_WARPS = 4;
constexpr int TB_SIZE = NUM_WARPS * WARP_SIZE;
constexpr int THREAD_M = 4;
__device__
void fp4x8_to_fp32x2x4(int in, int64_t *out) {
int tmp[4];
asm volatile(
"{\n"
".reg .b8 tmp0, tmp1, tmp2, tmp3;\n"
"mov.b32 {tmp0, tmp1, tmp2, tmp3}, %4; // unpack 32-bit register to 4x fp4x2\n"
"cvt.rn.f16x2.e2m1x2 %0, tmp0; // PTX only supports FP4->FP16\n"
"cvt.rn.f16x2.e2m1x2 %1, tmp1;\n"
"cvt.rn.f16x2.e2m1x2 %2, tmp2;\n"
"cvt.rn.f16x2.e2m1x2 %3, tmp3;\n"
"}\n"
: "=r"(tmp[0]), "=r"(tmp[1]), "=r"(tmp[2]), "=r"(tmp[3])
: "r"(in)
);
for (int i = 0; i < 4; i++)
asm volatile(
"{\n"
".reg .b16 b16_0, b16_1;\n"
".reg .b32 f32_0, f32_1;\n"
"mov.b32 {b16_0, b16_1}, %1; // unpack\n"
"cvt.f32.f16 f32_0, b16_0;\n"
"cvt.f32.f16 f32_1, b16_1;\n"
"mov.b64 %0, {f32_0, f32_1}; // pack\n"
"}\n"
: "=l"(out[i])
: "r"(tmp[i])
);
}
template <int HEIGHT, int WIDTH, int TB_SIZE, typename T>
__device__
void cp_async_2d(int dst, const T *src, int src_stride, int tid) {
auto load = [&](int idx) {
const int row = idx / WIDTH;
const int col = idx % WIDTH;
const int dst_addr = dst + idx * sizeof(T);
const T *src_addr = src + (row * src_stride + col);
asm volatile("cp.async.cg.shared.global [%0], [%1], 16;\n" :: "r"(dst_addr), "l"(src_addr));
};
constexpr int num_elems = 16 / sizeof(T);
constexpr int num_iters = HEIGHT * WIDTH / (TB_SIZE * num_elems);
for (int iter = 0; iter < num_iters; iter++)
load((iter * TB_SIZE + tid) * num_elems);
// handle the case when tile size is not divisible by threadblock size
if constexpr ((HEIGHT * WIDTH) % (TB_SIZE * num_elems) != 0) {
const int idx = (num_iters * TB_SIZE + tid) * num_elems;
if (idx < HEIGHT * WIDTH)
load(idx);
}
}
// to make our calculations simple, let's treat fp4x2 as a unit.
// hence, K = number of fp4x2 elements, and 8 elements share
// the same scale.
template <int THREAD_K, int NUM_STAGES = 1>
__global__
__launch_bounds__(NUM_WARPS * WARP_SIZE)
void kernel(
const char *A_ptr, // [L, M, K]
const char *B_ptr, // [L, 128, K]
const __nv_fp8_e4m3 *SFA_ptr, // [L, M, K/8]
const __nv_fp8_e4m3 *SFB_ptr, // [L, 128, K/8]
half *C_ptr, // [L, M]
int L, int M, int K
) {
// to ensure coalesced access, we need at least 8 threads per row (16B x 8 = 128B)
// each thread reads 16B, which covers 2 scaled groups. hence, we only need within
// thread reduction during the main loop.
static_assert(THREAD_M == 4);
static_assert(THREAD_K >= 8);
static_assert(THREAD_K <= TB_SIZE);
constexpr int BLOCK_M = (TB_SIZE / THREAD_K) * THREAD_M;
constexpr int BLOCK_K = THREAD_K * 16;
constexpr int SF_BLOCK_K = BLOCK_K / 8;
const int tid = threadIdx.x;
const int bid = blockIdx.x;
const int batch_id = blockIdx.y;
const int lane_id = tid % WARP_SIZE;
const int warp_id = tid / WARP_SIZE;
const int off_m = bid * BLOCK_M;
const int off_k = (tid % THREAD_K) * 16; // each thread reads 16 fp4x2 values at a time
A_ptr += (batch_id * M * K) + (off_m * K);
B_ptr += (batch_id * 128 * K);
SFA_ptr += (batch_id * M * (K / 8)) + (off_m * (K / 8));
SFB_ptr += (batch_id * 128 * (K / 8));
// set up smem
extern __shared__ char smem[];
const int smem_u32 = static_cast<int>(__cvta_generic_to_shared(smem));
constexpr int TOTAL_SMEM = (BLOCK_M * BLOCK_K) + (BLOCK_K) + (BLOCK_M * SF_BLOCK_K) + SF_BLOCK_K;
char *A_smem = smem;
char *B_smem = A_smem + BLOCK_M * BLOCK_K;
char *SFA_smem = B_smem + BLOCK_K;
char *SFB_smem = SFA_smem + BLOCK_M * SF_BLOCK_K;
// to be used for smem->rmem load
char *A_smem_ld = A_smem + (tid / THREAD_K) * THREAD_M * BLOCK_K + off_k;
char *B_smem_ld = B_smem + off_k;
char *SFA_smem_ld = SFA_smem + (tid / THREAD_K) * THREAD_M * SF_BLOCK_K + (off_k / 8);
char *SFB_smem_ld = SFB_smem + + (off_k / 8);
float acc[THREAD_M] = {};
auto load = [&](int iter_k) {
// NOTE: since B, SFA, and SFB does not require the whole threadblock to load, we can partition it within the threadblock.
const int buffer = smem_u32 + (iter_k % NUM_STAGES) * TOTAL_SMEM;
const int A_buf = buffer;
const int B_buf = A_buf + BLOCK_M * BLOCK_K;
const int SFA_buf = B_buf + BLOCK_K;
const int SFB_buf = SFA_buf + BLOCK_M * SF_BLOCK_K;
cp_async_2d<BLOCK_M, BLOCK_K, TB_SIZE>( A_buf, A_ptr, K, tid);
cp_async_2d< 1, BLOCK_K, TB_SIZE>( B_buf, B_ptr, K, tid);
cp_async_2d<BLOCK_M, SF_BLOCK_K, TB_SIZE>(SFA_buf, SFA_ptr, K / 8, tid);
cp_async_2d< 1, SF_BLOCK_K, TB_SIZE>(SFB_buf, SFB_ptr, K / 8, tid);
asm volatile("cp.async.commit_group;\n");
A_ptr += BLOCK_K;
B_ptr += BLOCK_K;
SFA_ptr += BLOCK_K / 8;
SFB_ptr += BLOCK_K / 8;
};
for (int iter_k = 0; iter_k < NUM_STAGES - 1; iter_k++)
load(iter_k);
const int num_iters = K / BLOCK_K;
for (int iter_k = 0; iter_k < num_iters; iter_k++) {
// gmem -> smem
if (iter_k + NUM_STAGES - 1 < num_iters) {
__syncthreads(); // make sure previous compute finish using the buffer
load(iter_k + NUM_STAGES - 1);
} else {
asm volatile("cp.async.commit_group;\n");
}
// smem -> rmem
asm volatile("cp.async.wait_group %0;\n" :: "n"(NUM_STAGES - 1));
__syncthreads(); // memory barrier
int A_fp4x8[THREAD_M][4], B_fp4x8[4];
float2 SFA_fp32x2[THREAD_M], SFB_fp32x2;
int buf_offset = (iter_k % NUM_STAGES) * TOTAL_SMEM;
for (int m = 0; m < THREAD_M; m++) {
reinterpret_cast<int4 *>(A_fp4x8[m])[0] = reinterpret_cast<const int4 *>(A_smem_ld + buf_offset + m * BLOCK_K)[0];
SFA_fp32x2[m] = static_cast<float2>(reinterpret_cast<const __nv_fp8x2_e4m3 *>(SFA_smem_ld + buf_offset + m * SF_BLOCK_K)[0]);
}
reinterpret_cast<int4 *>(B_fp4x8)[0] = reinterpret_cast<const int4 *>(B_smem_ld + buf_offset)[0];
SFB_fp32x2 = static_cast<float2>(reinterpret_cast<const __nv_fp8x2_e4m3 *>(SFB_smem_ld + buf_offset)[0]);
// unpack to FP32
int64_t A_fp32x2[THREAD_M][16], B_fp32x2[16];
for (int m = 0; m < THREAD_M; m++)
for (int i = 0; i < 4; i++)
fp4x8_to_fp32x2x4(A_fp4x8[m][i], A_fp32x2[m] + i * 4);
for (int i = 0; i < 4; i++)
fp4x8_to_fp32x2x4(B_fp4x8[i], B_fp32x2 + i * 4);
for (int m = 0; m < THREAD_M; m++)
for (int group_id = 0; group_id < 2; group_id++) {
// FMA. manually unroll the 1st iteration
int64_t sub_acc;
asm volatile("mul.rn.f32x2 %0, %1, %2;\n"
: "=l"(sub_acc)
: "l"(A_fp32x2[m][group_id * 8]), "l"(B_fp32x2[group_id * 8]));
for (int i = 1; i < 8; i++)
asm volatile("fma.rn.f32x2 %0, %1, %2, %0;\n"
: "+l"(sub_acc)
: "l"(A_fp32x2[m][group_id * 8 + i]), "l"(B_fp32x2[group_id * 8 + i]));
float tmp[2];
std::memcpy(tmp, &sub_acc, sizeof(sub_acc));
float sfa = reinterpret_cast<float *>(SFA_fp32x2 + m)[group_id];
float sfb = reinterpret_cast<float *>(&SFB_fp32x2)[group_id];
acc[m] += (tmp[0] + tmp[1]) * sfa * sfb;
}
}
// this is so cursed
long2 acc_fp32x2x2;
std::memcpy(&acc_fp32x2x2, acc, sizeof(acc_fp32x2x2));
// threadblock reduction
if constexpr (THREAD_K > WARP_SIZE) {
__shared__ long2 smem[TB_SIZE];
smem[tid] = acc_fp32x2x2;
__syncthreads();
for (int stride = THREAD_K / 2; stride >= WARP_SIZE; stride /= 2) {
if ((tid % THREAD_K) < stride) {
long2 tmp = smem[tid + stride];
asm volatile("add.rn.f32x2 %0, %0, %1;\n" : "+l"(acc_fp32x2x2.x) : "l"(tmp.x));
asm volatile("add.rn.f32x2 %0, %0, %1;\n" : "+l"(acc_fp32x2x2.y) : "l"(tmp.y));
smem[tid] = acc_fp32x2x2;
}
__syncthreads();
}
}
// warp reduction
constexpr int start_stride = std::min(THREAD_K, WARP_SIZE) / 2;
for (int stride = start_stride; stride > 0; stride /= 2) {
long tmp[2];
tmp[0] = __shfl_down_sync(0xFFFF'FFFF, acc_fp32x2x2.x, stride);
tmp[1] = __shfl_down_sync(0xFFFF'FFFF, acc_fp32x2x2.y, stride);
asm volatile("add.rn.f32x2 %0, %0, %1;\n" : "+l"(acc_fp32x2x2.x) : "l"(tmp[0]));
asm volatile("add.rn.f32x2 %0, %0, %1;\n" : "+l"(acc_fp32x2x2.y) : "l"(tmp[1]));
}
if (tid % THREAD_K == 0) {
half2 out[2];
out[0] = __float22half2_rn(reinterpret_cast<float2 *>(&acc_fp32x2x2)[0]);
out[1] = __float22half2_rn(reinterpret_cast<float2 *>(&acc_fp32x2x2)[1]);
reinterpret_cast<int2 *>(C_ptr + (batch_id * M + off_m + (tid / THREAD_K) * THREAD_M))[0] = reinterpret_cast<int2 *>(out)[0];
}
}
void gemv(
const at::Tensor& A,
const at::Tensor& B,
const at::Tensor& SFA,
const at::Tensor& SFB,
at::Tensor& C
) {
const int M = A.size(0);
const int K = A.size(1);
const int L = A.size(2);
auto A_ptr = reinterpret_cast<const char *>(A.data_ptr());
auto B_ptr = reinterpret_cast<const char *>(B.data_ptr());
auto SFA_ptr = reinterpret_cast<const __nv_fp8_e4m3 *>(SFA.data_ptr());
auto SFB_ptr = reinterpret_cast<const __nv_fp8_e4m3 *>(SFB.data_ptr());
auto C_ptr = reinterpret_cast<half *>(C.data_ptr());
auto stream = at::cuda::getCurrentCUDAStream();
constexpr int NUM_STAGES = 2;
#define launch(THREAD_K) { \
int BLOCK_M = (TB_SIZE / THREAD_K) * THREAD_M; \
int BLOCK_K = THREAD_K * 16; \
int SF_BLOCK_K = BLOCK_K / 8; \
int TOTAL_SMEM = (BLOCK_M * BLOCK_K) + (BLOCK_K) + (BLOCK_M * SF_BLOCK_K) + SF_BLOCK_K; \
dim3 grid(M / BLOCK_M, L); \
int smem_size = TOTAL_SMEM * NUM_STAGES; \
kernel<THREAD_K, NUM_STAGES><<<grid, TB_SIZE, smem_size, stream>>>(A_ptr, B_ptr, SFA_ptr, SFB_ptr, C_ptr, L, M, K); \
}
if (false) {}
else if (K % (128 * 16) == 0) launch(128) // benchmark.0
else if (K % (32 * 16) == 0) launch(32) // benchmark.1 and benchmark.2
else launch(8) // the rest
#undef launch
}
TORCH_LIBRARY(my_module, m) {
m.def("gemv(Tensor A, Tensor B, Tensor SFA, Tensor SFB, Tensor(a!) C) -> ()");
m.impl("gemv", &gemv);
}
"""
load_inline(
"gemv_c0",
cpp_sources="",
cuda_sources=CUDA_SRC,
verbose=True,
is_python_module=False,
no_implicit_headers=True,
extra_cuda_cflags=[
"-O3",
"-gencode=arch=compute_100a,code=sm_100a",
"-gencode=arch=compute_120a,code=sm_120a",
"-lineinfo",
],
)
def custom_kernel(data: input_t) -> output_t:
# a: [ M, K, L], natural shape [L, M, K]
# b: [128, K, L], natural shape [L, 128, K] - only the 1st row is used
# sfa: [32, 4, rest_m, 4, rest_k, L], natural shape [L, rest_m, rest_k, 32, 4, 4]
# sfb: [32, 4, 1, 4, rest_k, L], natural shape [L, 1, rest_k, 32, 4, 4]
# c: [ M, 1, L], natural shape [L, M, 1]
a, b, sfa, sfb, _, _, c_ref = data
torch.ops.my_module.gemv(a, b, sfa, sfb, c_ref)
if False:
M, K, L = a.shape
path = Path(f"profile_data/{M=}_K={K * 2}_{L=}.json.gz")
if not path.exists():
a.new_zeros(int(1e8), dtype=torch.uint8) # 100 MB
with torch.profiler.profile() as prof:
torch.ops.my_module.gemv(a, b, sfa, sfb, c_ref)
path.parent.mkdir(exist_ok=True)
prof.export_chrome_trace(str(path))
return c_ref
scrolls · 342 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 73354.
⋯ 5 unchanged linesfrom task import input_t, output_tfrom torch.utils.cpp_extension import load_inline- # https://github.com/NVIDIA/cutlass/blob/v4.2.1/examples/72_blackwell_narrow_precision_gemm/72b_blackwell_nvfp4_nvfp4_gemm.cuCUDA_SRC = r"""- #include "cutlass/cutlass.h"+ #include <cuda_fp16.h>+ #include <cuda_fp8.h>- #include "cute/tensor.hpp"- #include "cutlass/tensor_ref.h"- #include "cutlass/epilogue/thread/linear_combination.h"- #include "cutlass/gemm/dispatch_policy.hpp"- #include "cutlass/gemm/collective/collective_builder.hpp"- #include "cutlass/epilogue/collective/collective_builder.hpp"- #include "cutlass/detail/sm100_blockscaled_layout.hpp"- #include "cutlass/gemm/device/gemm_universal_adapter.h"- #include "cutlass/gemm/kernel/gemm_universal.hpp"- #include "cutlass/gemm/kernel/tile_scheduler_params.h"-- #include "cutlass/util/packed_stride.hpp"-#include <torch/library.h>#include <ATen/ATen.h>#include <ATen/core/Tensor.h>#include <ATen/cuda/CUDAUtils.h>#include <ATen/cuda/CUDAContext.h>- #define STRINGIFY(x) #x- #define CUTLASS_CHECK(call) \- do { \- auto status = call; \- TORCH_CHECK(status == cutlass::Status::kSuccess, STRINGIFY(call), ": ", status, " - ", cutlassGetStatusString(status)); \- } while (0)+ constexpr int WARP_SIZE = 32;+ constexpr int NUM_WARPS = 4;+ constexpr int TB_SIZE = NUM_WARPS * WARP_SIZE;+ constexpr int THREAD_M = 4;- using namespace cute;+ __device__+ void fp4x8_to_fp32x2x4(int in, int64_t *out) {+ int tmp[4];+ asm volatile(+ "{\n"+ ".reg .b8 tmp0, tmp1, tmp2, tmp3;\n"+ "mov.b32 {tmp0, tmp1, tmp2, tmp3}, %4; // unpack 32-bit register to 4x fp4x2\n"+ "cvt.rn.f16x2.e2m1x2 %0, tmp0; // PTX only supports FP4->FP16\n"+ "cvt.rn.f16x2.e2m1x2 %1, tmp1;\n"+ "cvt.rn.f16x2.e2m1x2 %2, tmp2;\n"+ "cvt.rn.f16x2.e2m1x2 %3, tmp3;\n"+ "}\n"+ : "=r"(tmp[0]), "=r"(tmp[1]), "=r"(tmp[2]), "=r"(tmp[3])+ : "r"(in)+ );- using ElementAB = cutlass::nv_float4_t<cutlass::float_e2m1_t>;- using ElementC = cutlass::half_t;- using ElementAcc = float;+ for (int i = 0; i < 4; i++)+ asm volatile(+ "{\n"+ ".reg .b16 b16_0, b16_1;\n"+ ".reg .b32 f32_0, f32_1;\n"+ "mov.b32 {b16_0, b16_1}, %1; // unpack\n"+ "cvt.f32.f16 f32_0, b16_0;\n"+ "cvt.f32.f16 f32_1, b16_1;\n"+ "mov.b64 %0, {f32_0, f32_1}; // pack\n"+ "}\n"+ : "=l"(out[i])+ : "r"(tmp[i])+ );+ }- constexpr int AlignmentAB = 128 / 4; // 32- constexpr int AlignmentC = 128 / cutlass::sizeof_bits<ElementC>::value; // 8+ template <int HEIGHT, int WIDTH, int TB_SIZE, typename T>+ __device__+ void cp_async_2d(int dst, const T *src, int src_stride, int tid) {+ auto load = [&](int idx) {+ const int row = idx / WIDTH;+ const int col = idx % WIDTH;- using LayoutATag = cutlass::layout::RowMajor;- using LayoutBTag = cutlass::layout::ColumnMajor;- using LayoutCTag = cutlass::layout::RowMajor;+ const int dst_addr = dst + idx * sizeof(T);+ const T *src_addr = src + (row * src_stride + col);+ asm volatile("cp.async.cg.shared.global [%0], [%1], 16;\n" :: "r"(dst_addr), "l"(src_addr));+ };- using ArchTag = cutlass::arch::Sm100;- using OperatorClass = cutlass::arch::OpClassBlockScaledTensorOp;+ constexpr int num_elems = 16 / sizeof(T);+ constexpr int num_iters = HEIGHT * WIDTH / (TB_SIZE * num_elems);- // Kernel Perf config- using MmaTileShape = Shape<_128,_128,_256>;- using ClusterShape = Shape<_1,_1,_1>;+ for (int iter = 0; iter < num_iters; iter++)+ load((iter * TB_SIZE + tid) * num_elems);- using CollectiveEpilogue = typename cutlass::epilogue::collective::CollectiveBuilder<- ArchTag, OperatorClass,- MmaTileShape, ClusterShape,- cutlass::epilogue::collective::EpilogueTileAuto,- ElementAcc, ElementAcc,- ElementC, LayoutCTag, AlignmentC,- ElementC, LayoutCTag, AlignmentC,- cutlass::epilogue::collective::EpilogueScheduleAuto- >::CollectiveOp;+ // handle the case when tile size is not divisible by threadblock size+ if constexpr ((HEIGHT * WIDTH) % (TB_SIZE * num_elems) != 0) {+ const int idx = (num_iters * TB_SIZE + tid) * num_elems;+ if (idx < HEIGHT * WIDTH)+ load(idx);+ }+ }- using CollectiveMainloop = typename cutlass::gemm::collective::CollectiveBuilder<- ArchTag, OperatorClass,- ElementAB, LayoutATag, AlignmentAB,- ElementAB, LayoutBTag, AlignmentAB,- ElementAcc,- MmaTileShape, ClusterShape,- cutlass::gemm::collective::StageCountAutoCarveout<static_cast<int>(sizeof(typename CollectiveEpilogue::SharedStorage))>,- cutlass::gemm::collective::KernelScheduleAuto- >::CollectiveOp;+ // to make our calculations simple, let's treat fp4x2 as a unit.+ // hence, K = number of fp4x2 elements, and 8 elements share+ // the same scale.+ template <int THREAD_K, int NUM_STAGES = 1>+ __global__+ __launch_bounds__(NUM_WARPS * WARP_SIZE)+ void kernel(+ const char *A_ptr, // [L, M, K]+ const char *B_ptr, // [L, 128, K]+ const __nv_fp8_e4m3 *SFA_ptr, // [L, M, K/8]+ const __nv_fp8_e4m3 *SFB_ptr, // [L, 128, K/8]+ half *C_ptr, // [L, M]+ int L, int M, int K+ ) {+ // to ensure coalesced access, we need at least 8 threads per row (16B x 8 = 128B)+ // each thread reads 16B, which covers 2 scaled groups. hence, we only need within+ // thread reduction during the main loop.+ static_assert(THREAD_M == 4);+ static_assert(THREAD_K >= 8);+ static_assert(THREAD_K <= TB_SIZE);+ constexpr int BLOCK_M = (TB_SIZE / THREAD_K) * THREAD_M;+ constexpr int BLOCK_K = THREAD_K * 16;+ constexpr int SF_BLOCK_K = BLOCK_K / 8;- using GemmKernel = cutlass::gemm::kernel::GemmUniversal<- Shape<int, int, int, int>,- CollectiveMainloop,- CollectiveEpilogue,- void>;+ const int tid = threadIdx.x;+ const int bid = blockIdx.x;+ const int batch_id = blockIdx.y;- using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;+ const int lane_id = tid % WARP_SIZE;+ const int warp_id = tid / WARP_SIZE;+ const int off_m = bid * BLOCK_M;+ const int off_k = (tid % THREAD_K) * 16; // each thread reads 16 fp4x2 values at a time++ A_ptr += (batch_id * M * K) + (off_m * K);+ B_ptr += (batch_id * 128 * K);++ SFA_ptr += (batch_id * M * (K / 8)) + (off_m * (K / 8));+ SFB_ptr += (batch_id * 128 * (K / 8));++ // set up smem+ extern __shared__ char smem[];+ const int smem_u32 = static_cast<int>(__cvta_generic_to_shared(smem));+ constexpr int TOTAL_SMEM = (BLOCK_M * BLOCK_K) + (BLOCK_K) + (BLOCK_M * SF_BLOCK_K) + SF_BLOCK_K;++ char *A_smem = smem;+ char *B_smem = A_smem + BLOCK_M * BLOCK_K;+ char *SFA_smem = B_smem + BLOCK_K;+ char *SFB_smem = SFA_smem + BLOCK_M * SF_BLOCK_K;++ // to be used for smem->rmem load+ char *A_smem_ld = A_smem + (tid / THREAD_K) * THREAD_M * BLOCK_K + off_k;+ char *B_smem_ld = B_smem + off_k;+ char *SFA_smem_ld = SFA_smem + (tid / THREAD_K) * THREAD_M * SF_BLOCK_K + (off_k / 8);+ char *SFB_smem_ld = SFB_smem + + (off_k / 8);++ float acc[THREAD_M] = {};++ auto load = [&](int iter_k) {+ // NOTE: since B, SFA, and SFB does not require the whole threadblock to load, we can partition it within the threadblock.+ const int buffer = smem_u32 + (iter_k % NUM_STAGES) * TOTAL_SMEM;+ const int A_buf = buffer;+ const int B_buf = A_buf + BLOCK_M * BLOCK_K;+ const int SFA_buf = B_buf + BLOCK_K;+ const int SFB_buf = SFA_buf + BLOCK_M * SF_BLOCK_K;++ cp_async_2d<BLOCK_M, BLOCK_K, TB_SIZE>( A_buf, A_ptr, K, tid);+ cp_async_2d< 1, BLOCK_K, TB_SIZE>( B_buf, B_ptr, K, tid);+ cp_async_2d<BLOCK_M, SF_BLOCK_K, TB_SIZE>(SFA_buf, SFA_ptr, K / 8, tid);+ cp_async_2d< 1, SF_BLOCK_K, TB_SIZE>(SFB_buf, SFB_ptr, K / 8, tid);++ asm volatile("cp.async.commit_group;\n");++ A_ptr += BLOCK_K;+ B_ptr += BLOCK_K;+ SFA_ptr += BLOCK_K / 8;+ SFB_ptr += BLOCK_K / 8;+ };++ for (int iter_k = 0; iter_k < NUM_STAGES - 1; iter_k++)+ load(iter_k);++ const int num_iters = K / BLOCK_K;++ for (int iter_k = 0; iter_k < num_iters; iter_k++) {+ // gmem -> smem+ if (iter_k + NUM_STAGES - 1 < num_iters) {+ __syncthreads(); // make sure previous compute finish using the buffer+ load(iter_k + NUM_STAGES - 1);+ } else {+ asm volatile("cp.async.commit_group;\n");+ }++ // smem -> rmem+ asm volatile("cp.async.wait_group %0;\n" :: "n"(NUM_STAGES - 1));+ __syncthreads(); // memory barrier++ int A_fp4x8[THREAD_M][4], B_fp4x8[4];+ float2 SFA_fp32x2[THREAD_M], SFB_fp32x2;+ int buf_offset = (iter_k % NUM_STAGES) * TOTAL_SMEM;++ for (int m = 0; m < THREAD_M; m++) {+ reinterpret_cast<int4 *>(A_fp4x8[m])[0] = reinterpret_cast<const int4 *>(A_smem_ld + buf_offset + m * BLOCK_K)[0];+ SFA_fp32x2[m] = static_cast<float2>(reinterpret_cast<const __nv_fp8x2_e4m3 *>(SFA_smem_ld + buf_offset + m * SF_BLOCK_K)[0]);+ }++ reinterpret_cast<int4 *>(B_fp4x8)[0] = reinterpret_cast<const int4 *>(B_smem_ld + buf_offset)[0];+ SFB_fp32x2 = static_cast<float2>(reinterpret_cast<const __nv_fp8x2_e4m3 *>(SFB_smem_ld + buf_offset)[0]);++ // unpack to FP32+ int64_t A_fp32x2[THREAD_M][16], B_fp32x2[16];++ for (int m = 0; m < THREAD_M; m++)+ for (int i = 0; i < 4; i++)+ fp4x8_to_fp32x2x4(A_fp4x8[m][i], A_fp32x2[m] + i * 4);++ for (int i = 0; i < 4; i++)+ fp4x8_to_fp32x2x4(B_fp4x8[i], B_fp32x2 + i * 4);++ for (int m = 0; m < THREAD_M; m++)+ for (int group_id = 0; group_id < 2; group_id++) {+ // FMA. manually unroll the 1st iteration+ int64_t sub_acc;+ asm volatile("mul.rn.f32x2 %0, %1, %2;\n"+ : "=l"(sub_acc)+ : "l"(A_fp32x2[m][group_id * 8]), "l"(B_fp32x2[group_id * 8]));+ for (int i = 1; i < 8; i++)+ asm volatile("fma.rn.f32x2 %0, %1, %2, %0;\n"+ : "+l"(sub_acc)+ : "l"(A_fp32x2[m][group_id * 8 + i]), "l"(B_fp32x2[group_id * 8 + i]));++ float tmp[2];+ std::memcpy(tmp, &sub_acc, sizeof(sub_acc));++ float sfa = reinterpret_cast<float *>(SFA_fp32x2 + m)[group_id];+ float sfb = reinterpret_cast<float *>(&SFB_fp32x2)[group_id];+ acc[m] += (tmp[0] + tmp[1]) * sfa * sfb;+ }+ }++ // this is so cursed+ long2 acc_fp32x2x2;+ std::memcpy(&acc_fp32x2x2, acc, sizeof(acc_fp32x2x2));++ // threadblock reduction+ if constexpr (THREAD_K > WARP_SIZE) {+ __shared__ long2 smem[TB_SIZE];+ smem[tid] = acc_fp32x2x2;+ __syncthreads();++ for (int stride = THREAD_K / 2; stride >= WARP_SIZE; stride /= 2) {+ if ((tid % THREAD_K) < stride) {+ long2 tmp = smem[tid + stride];+ asm volatile("add.rn.f32x2 %0, %0, %1;\n" : "+l"(acc_fp32x2x2.x) : "l"(tmp.x));+ asm volatile("add.rn.f32x2 %0, %0, %1;\n" : "+l"(acc_fp32x2x2.y) : "l"(tmp.y));+ smem[tid] = acc_fp32x2x2;+ }+ __syncthreads();+ }+ }++ // warp reduction+ constexpr int start_stride = std::min(THREAD_K, WARP_SIZE) / 2;+ for (int stride = start_stride; stride > 0; stride /= 2) {+ long tmp[2];+ tmp[0] = __shfl_down_sync(0xFFFF'FFFF, acc_fp32x2x2.x, stride);+ tmp[1] = __shfl_down_sync(0xFFFF'FFFF, acc_fp32x2x2.y, stride);+ asm volatile("add.rn.f32x2 %0, %0, %1;\n" : "+l"(acc_fp32x2x2.x) : "l"(tmp[0]));+ asm volatile("add.rn.f32x2 %0, %0, %1;\n" : "+l"(acc_fp32x2x2.y) : "l"(tmp[1]));+ }++ if (tid % THREAD_K == 0) {+ half2 out[2];+ out[0] = __float22half2_rn(reinterpret_cast<float2 *>(&acc_fp32x2x2)[0]);+ out[1] = __float22half2_rn(reinterpret_cast<float2 *>(&acc_fp32x2x2)[1]);+ reinterpret_cast<int2 *>(C_ptr + (batch_id * M + off_m + (tid / THREAD_K) * THREAD_M))[0] = reinterpret_cast<int2 *>(out)[0];+ }+ }+void gemv(const at::Tensor& A,const at::Tensor& B,⋯ 2 unchanged linesat::Tensor& C) {const int M = A.size(0);- const int N = 128;- const int K = A.size(1) * 2;+ const int K = A.size(1);const int L = A.size(2);- using ABType = typename ElementAB::DataType;- using SFType = typename ElementAB::ScaleFactorType;+ auto A_ptr = reinterpret_cast<const char *>(A.data_ptr());+ auto B_ptr = reinterpret_cast<const char *>(B.data_ptr());+ auto SFA_ptr = reinterpret_cast<const __nv_fp8_e4m3 *>(SFA.data_ptr());+ auto SFB_ptr = reinterpret_cast<const __nv_fp8_e4m3 *>(SFB.data_ptr());+ auto C_ptr = reinterpret_cast<half *>(C.data_ptr());- auto stride_A = cutlass::make_cute_packed_stride(typename GemmKernel::StrideA{}, {M, K, L});- auto stride_B = cutlass::make_cute_packed_stride(typename GemmKernel::StrideB{}, {N, K, L});- auto stride_C = cutlass::make_cute_packed_stride(typename GemmKernel::StrideC{}, {M, N, L});+ auto stream = at::cuda::getCurrentCUDAStream();+ constexpr int NUM_STAGES = 2;- using Sm1xxBlkScaledConfig = typename Gemm::GemmKernel::CollectiveMainloop::Sm1xxBlkScaledConfig;- auto layout_SFA = Sm1xxBlkScaledConfig::tile_atom_to_shape_SFA(cute::make_shape(M, N, K, L));- auto layout_SFB = Sm1xxBlkScaledConfig::tile_atom_to_shape_SFB(cute::make_shape(M, N, K, L));+ #define launch(THREAD_K) { \+ int BLOCK_M = (TB_SIZE / THREAD_K) * THREAD_M; \+ int BLOCK_K = THREAD_K * 16; \+ int SF_BLOCK_K = BLOCK_K / 8; \+ int TOTAL_SMEM = (BLOCK_M * BLOCK_K) + (BLOCK_K) + (BLOCK_M * SF_BLOCK_K) + SF_BLOCK_K; \+ dim3 grid(M / BLOCK_M, L); \+ int smem_size = TOTAL_SMEM * NUM_STAGES; \+ kernel<THREAD_K, NUM_STAGES><<<grid, TB_SIZE, smem_size, stream>>>(A_ptr, B_ptr, SFA_ptr, SFB_ptr, C_ptr, L, M, K); \+ }- auto *A_ptr = reinterpret_cast<const ABType *>(A.data_ptr());- auto *B_ptr = reinterpret_cast<const ABType *>(B.data_ptr());- auto *SFA_ptr = reinterpret_cast<const SFType *>(SFA.data_ptr());- auto *SFB_ptr = reinterpret_cast<const SFType *>(SFB.data_ptr());- auto *C_ptr = reinterpret_cast<ElementC *>(C.data_ptr());+ if (false) {}+ else if (K % (128 * 16) == 0) launch(128) // benchmark.0+ else if (K % (32 * 16) == 0) launch(32) // benchmark.1 and benchmark.2+ else launch(8) // the rest- typename Gemm::Arguments arguments{- cutlass::gemm::GemmUniversalMode::kGemm,- {M, N, K, L},- {- A_ptr, stride_A,- B_ptr, stride_B,- SFA_ptr, layout_SFA,- SFB_ptr, layout_SFB,- },- {- {1.0f, 0.0f}, // alpha and beta- C_ptr, stride_C,- C_ptr, stride_C,- }- };-- Gemm gemm;- //CUTLASS_CHECK(gemm.can_implement(arguments));-- //long workspace_size = Gemm::get_workspace_size(arguments);- //at::Tensor workspace = at::empty({workspace_size}, A.options().dtype(at::kByte));- auto stream = at::cuda::getCurrentCUDAStream();-- //CUTLASS_CHECK(gemm.initialize(arguments, workspace.data_ptr(), stream));- CUTLASS_CHECK(gemm.initialize(arguments, 0, stream));- CUTLASS_CHECK(gemm.run(stream));+ #undef launch}TORCH_LIBRARY(my_module, m) {⋯ 12 unchanged linesextra_cuda_cflags=["-O3","-gencode=arch=compute_100a,code=sm_100a",+ "-gencode=arch=compute_120a,code=sm_120a",+ "-lineinfo",],)⋯ 4 unchanged lines# sfa: [32, 4, rest_m, 4, rest_k, L], natural shape [L, rest_m, rest_k, 32, 4, 4]# sfb: [32, 4, 1, 4, rest_k, L], natural shape [L, 1, rest_k, 32, 4, 4]# c: [ M, 1, L], natural shape [L, M, 1]- a, b, _, _, sfa, sfb, c_ref = data+ a, b, sfa, sfb, _, _, c_ref = data+ torch.ops.my_module.gemv(a, b, sfa, sfb, c_ref)- M = a.shape[0]- N = 128- K = a.shape[1] * 2- L = a.shape[2]-- big_c = c_ref.new_empty(L, M, N)- torch.ops.my_module.gemv(a, b, sfa, sfb, big_c)-if False:- path = Path(f"profile_data/{M=}_{K=}_{L=}.json.gz")+ M, K, L = a.shape+ path = Path(f"profile_data/{M=}_K={K * 2}_{L=}.json.gz")if not path.exists():a.new_zeros(int(1e8), dtype=torch.uint8) # 100 MBwith torch.profiler.profile() as prof:- torch.ops.my_module.gemv(a, b, sfa, sfb, big_c)+ torch.ops.my_module.gemv(a, b, sfa, sfb, c_ref)path.parent.mkdir(exist_ok=True)prof.export_chrome_trace(str(path))- return big_c[..., :1].permute(1, 2, 0) # convert to [M, 1, L]+ return c_ref
scrolls · 434 diff lines total
Best evidence level for this revision: reported
JSON