submission 297876
symlon · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 574 lines, June 9 Researcher Reciprocity License v1.0.
inline.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-nvfp4-dual-gemm-297876?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:26c39caab57dbc604d05a22222417fc07feb52760b44859ddf17b26b8e082742
license declaredunknown
license concludedunknown
authorssymlon
imported2026-08-26
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
fp4
PyTorch reference implementation of NVFP4 block-scaled dual GEMM with silu activation,fp8
using sfA_gl = gl<fp8e4m3, -1, -1, PACKED_PER_TILE * 32, 16, sfA_tile>;shared-memory
extern __shared__ int __shm[];stages = 2
static constexpr int PIPELINE_STAGES = 2;warp-specialization
static constexpr int CONSUMER_WARPGROUPS = 1;Kernel source
inline.py574 lines
#!POPCORN leaderboard nvfp4_dual_gemm
#!POPCORN gpu NVIDIA
import os
import subprocess
import shutil
import torch
from torch.utils.cpp_extension import load_inline
from task import input_t, output_t
thunderkittens_dir = "/home/runner/ThunderKittens"
base_path = "/home/runner"
dirs = [d for d in os.listdir(base_path) if os.path.isdir(os.path.join(base_path, d))]
print("Existing directories in /home/runner:")
for d in dirs:
print(d)
if os.path.exists(thunderkittens_dir):
shutil.rmtree(thunderkittens_dir)
print(f"Deleted existing folder {thunderkittens_dir}.")
if not os.path.exists(thunderkittens_dir):
result = subprocess.run([
"git", "clone", "--recursive",
"-b", "tk-v2",
"https://github.com/symlons/ThunderKittens.git",
thunderkittens_dir
])
if result.returncode == 0:
print(f"ThunderKittens cloned successfully into {thunderkittens_dir}.")
else:
print("Failed to clone ThunderKittens.")
else:
print(f"{thunderkittens_dir} already exists. Switching to tk-v2 branch.")
subprocess.run(["git", "-C", thunderkittens_dir, "fetch", "origin"])
subprocess.run(["git", "-C", thunderkittens_dir, "checkout", "tk-v2"])
subprocess.run(["git", "-C", thunderkittens_dir, "pull"], check=True)
subprocess.run(["git", "-C", thunderkittens_dir, "reset", "--hard", "origin/tk-v2"])
subprocess.run(["git", "-C", thunderkittens_dir, "log", "-1", "--format=%H"])
nvpf4_src = r'''
#include "kittens.cuh"
#include "pyutils/torchutils.cuh"
// #include "ATen/ops/linalg_vector_norm.h"
// #include "ATen/ops/unsqueeze.h"
// todo: 2-cta pipeline
// todo: work-stealing
// todo: register allocation
// todo: warp index order?
// todo: tuning
// todo: ncu profile
using namespace kittens;
namespace nvfp4_gemm {
struct config {
static constexpr int CLUSTER_SIZE = 1;
static constexpr int NUM_BLOCKS = 148;
static constexpr int STATIC_SHARED_MEMORY = 1024;
static constexpr int DYNAMIC_SHARED_MEMORY = MAX_SHARED_MEMORY - STATIC_SHARED_MEMORY;
static constexpr int CONSUMER_WARPGROUPS = 1;
static constexpr int PRODUCER_WARPGROUPS = 1;
static constexpr int NUM_WARPGROUPS = CONSUMER_WARPGROUPS + PRODUCER_WARPGROUPS;
static constexpr int NUM_WARPS = NUM_WARPGROUPS * WARPGROUP_WARPS;
static constexpr int NUM_THREADS = NUM_WARPS * WARP_THREADS;
static constexpr int PRODUCER_REGISTERS = 256;
static constexpr int CONSUMER_REGISTERS = 256;
};
struct globals {
static constexpr int PIPELINE_STAGES = 2;
static constexpr int TILE_SIZE = 128; // todo: up to 256
static constexpr int PACKED_PER_TILE = 4;
static constexpr int REDUCTION_BLOCK = PACKED_PER_TILE * 32;
using A_fp4x2_tile = st_fp4e2m1_2<TILE_SIZE, TILE_SIZE>;
using A_fp4x2_gl = gl<fp4e2m1_2, 1, 1, -1, -1, A_fp4x2_tile>;
using B_fp4x2_tile = st_fp4e2m1_2<TILE_SIZE, TILE_SIZE>;
using B_fp4x2_gl = gl<fp4e2m1_2, 1, 1, -1, -1, B_fp4x2_tile>;
using C_tile = st_hf<TILE_SIZE, TILE_SIZE>;
using C_gl = gl<half, 1, 1, -1, -1, C_tile>;
A_fp4x2_gl A;
B_fp4x2_gl B;
B_fp4x2_gl B_2;
C_gl C;
using sfA_tile = st_fp8e4m3<PACKED_PER_TILE * 32, 16, false>;
using sfA_gl = gl<fp8e4m3, -1, -1, PACKED_PER_TILE * 32, 16, sfA_tile>;
using sfB_tile = st_fp8e4m3<PACKED_PER_TILE * 32, 16, false>;
using sfB_gl = gl<fp8e4m3, -1, -1, PACKED_PER_TILE * 32, 16, sfB_tile>;
sfA_gl sfA;
sfB_gl sfB;
sfB_gl sfB_2;
// __host__ inline dim3 grid() const {
// return dim3(A.cols() / TILE_SIZE, A.rows() / TILE_SIZE);
// }
__host__ inline int dynamic_shared_memory() const { // todo: check this
return 4 * TILE_SIZE * TILE_SIZE * sizeof(bf16) + 1024;
}
};
struct pipeline_input_tiles {
globals::A_fp4x2_tile A;
globals::B_fp4x2_tile B;
globals::B_fp4x2_tile B_2;
};
struct pipeline_input_scales {
globals::sfA_tile sfA;
globals::sfB_tile sfB;
globals::sfB_tile sfB_2;
};
struct pipeline_outputs {
globals::C_tile C;
};
__device__ inline void kernel(const globals &G) {
extern __shared__ int __shm[];
tma_swizzle_allocator sm_allocator((int *)&__shm[0]);
pipeline_input_tiles(&input_tiles)[globals::PIPELINE_STAGES] = sm_allocator.allocate<pipeline_input_tiles, globals::PIPELINE_STAGES>();
pipeline_input_scales(&input_scales)[globals::PIPELINE_STAGES] = sm_allocator.allocate<pipeline_input_scales, globals::PIPELINE_STAGES>();
pipeline_outputs &output_tiles = sm_allocator.allocate<pipeline_outputs>();
tensor_allocator<1, config::CLUSTER_SIZE> tm_allocator;
auto out_tm = tm_allocator.allocate<full_tt_fl<128>>(0);
auto out_tm_2 = tm_allocator.allocate<full_tt_fl<128>>(128);
auto A_sc_tm = tm_allocator.allocate<full_tt_fp8e4m3<16 * globals::PACKED_PER_TILE * globals::PIPELINE_STAGES>>(256);
auto B_sc_tm = tm_allocator.allocate<full_tt_fp8e4m3<16 * globals::PACKED_PER_TILE * globals::PIPELINE_STAGES>>(256 + 4 * globals::PACKED_PER_TILE * globals::PIPELINE_STAGES);
auto B_sc_tm_2 = tm_allocator.allocate<full_tt_fp8e4m3<16 * globals::PACKED_PER_TILE * globals::PIPELINE_STAGES>>(256 + 8 * globals::PACKED_PER_TILE * globals::PIPELINE_STAGES);
__shared__ semaphore inputs_smem_arrived[globals::PIPELINE_STAGES];
__shared__ semaphore matmul_finished_1[globals::PIPELINE_STAGES];
__shared__ semaphore matmul_finished_2[globals::PIPELINE_STAGES];
__shared__ semaphore inputs_all_arrived[globals::PIPELINE_STAGES];
__shared__ semaphore outputs_arrived;
__shared__ semaphore tensor_finished;
if (threadIdx.x == 32) {
#pragma unroll
for (int i = 0; i < globals::PIPELINE_STAGES; ++i) {
init_semaphore(matmul_finished_1[i], 0, 1);
init_semaphore(matmul_finished_2[i], 0, 1);
init_semaphore(inputs_all_arrived[i], 0, 1);
init_semaphore(inputs_smem_arrived[i], 0, 1);
}
init_semaphore(outputs_arrived, 0, 1);
init_semaphore(tensor_finished, 0, 1);
}
everyone::tma::cluster::sync();
int lane_id = warp::laneid();
int warp_id = warpgroup::warpid();
int warpgroup_id = warpgroup::groupid();
int cta_id = cluster_ctarank();
int cluster_id = clusterIdx().x;
// todo: super blocking
int num_blocks_per_col = G.C.cols() / globals::TILE_SIZE;
int num_blocks_per_row = G.C.rows() / globals::TILE_SIZE;
int num_blocks = num_blocks_per_row * num_blocks_per_col;
int num_iters_per_block = G.A.cols() / globals::REDUCTION_BLOCK; // todo: fix
uint32_t stage = 0;
uint32_t phasebits = 0xFFFF0000;
uint32_t phasebits_2 = 0xFFFF0000;
uint32_t last_stage = globals::PIPELINE_STAGES;
__syncthreads();
if (warpgroup_id == config::NUM_WARPGROUPS - 1) { // Producer group
// warpgroup::increase_registers<config::PRODUCER_REGISTERS>();
// warpgroup::decrease_registers<config::PRODUCER_REGISTERS>();
if (warp_id == 1 && lane_id == 0) {
// printf("G.A.cols() = %d\n", G.A.cols());
// printf("num_iters_per_block = %d\n", num_iters_per_block);
for (int block_idx = cluster_id; block_idx < num_blocks; block_idx += gridDim.x / config::CLUSTER_SIZE) {
int row_block_idx = block_idx / num_blocks_per_col;
int col_block_idx = block_idx % num_blocks_per_col;
for (int i = 0; i < num_iters_per_block; ++i) {
wait(matmul_finished_1[stage], get_phasebit<1>(phasebits, stage));
wait(matmul_finished_2[stage], get_phasebit<1>(phasebits_2, stage));
update_phasebit<1>(phasebits, stage);
update_phasebit<1>(phasebits_2, stage);
if (stage == last_stage) {
arrive(outputs_arrived);
last_stage = globals::PIPELINE_STAGES;
}
tma::expect_bytes(inputs_smem_arrived[stage], sizeof(globals::A_fp4x2_tile) + sizeof(globals::sfA_tile) + (sizeof(globals::B_fp4x2_tile) + sizeof(globals::sfB_tile)) * 2);
tma::load_async(input_tiles[stage].A, G.A, {row_block_idx, i}, inputs_smem_arrived[stage]);
tma::load_async(input_tiles[stage].B, G.B, {col_block_idx, i}, inputs_smem_arrived[stage]);
tma::load_async(input_tiles[stage].B_2, G.B_2, {col_block_idx, i}, inputs_smem_arrived[stage]);
// todo: consider sepearete loading warp for B2
tma::load_async(input_scales[stage].sfA, G.sfA, {row_block_idx, i, 0, 0}, inputs_smem_arrived[stage]);
tma::load_async(input_scales[stage].sfB, G.sfB, {col_block_idx, i, 0, 0}, inputs_smem_arrived[stage]);
tma::load_async(input_scales[stage].sfB_2, G.sfB_2, {col_block_idx, i, 0, 0}, inputs_smem_arrived[stage]);
if (i == num_iters_per_block - 1) {
last_stage = stage;
}
stage = (stage + 1) % globals::PIPELINE_STAGES;
}
}
if (last_stage < globals::PIPELINE_STAGES) {
wait(matmul_finished_1[last_stage], get_phasebit<1>(phasebits, last_stage));
wait(matmul_finished_2[last_stage], get_phasebit<1>(phasebits_2, last_stage));
arrive(outputs_arrived);
}
} else if (warp_id == 0 && lane_id == 0) {
for (int block_idx = cluster_id; block_idx < num_blocks; block_idx += gridDim.x / config::CLUSTER_SIZE) {
for (int i = 0; i < num_iters_per_block; i++) {
wait(inputs_smem_arrived[stage], get_phasebit<0>(phasebits, stage));
wait(inputs_smem_arrived[stage], get_phasebit<0>(phasebits_2, stage));
update_phasebit<0>(phasebits, stage);
update_phasebit<0>(phasebits_2, stage);
#pragma unroll
for (int ii = 0; ii < globals::PACKED_PER_TILE; ii++) {
auto A_sc_tm_subtile = A_sc_tm.subtile<full_tt_fp8e4m3<16>>(stage * globals::PACKED_PER_TILE * 16 + ii * 16);
auto B_sc_tm_subtile = B_sc_tm.subtile<full_tt_fp8e4m3<16>>(stage * globals::PACKED_PER_TILE * 16 + ii * 16);
auto &A_sc_sm_subtile = *reinterpret_cast<st_fp8e4m3<32, 16, false> *>(reinterpret_cast<uint64_t>(&input_scales[stage].sfA.data[0]) + 16 * 32 * ii);
auto &B_sc_sm_subtile = *reinterpret_cast<st_fp8e4m3<32, 16, false> *>(reinterpret_cast<uint64_t>(&input_scales[stage].sfB.data[0]) + 16 * 32 * ii);
load_mxnv_scale_async(A_sc_tm_subtile, A_sc_sm_subtile);
load_mxnv_scale_async(B_sc_tm_subtile, B_sc_sm_subtile);
auto B_sc_tm_subtile_2 = B_sc_tm_2.subtile<full_tt_fp8e4m3<16>>(stage * globals::PACKED_PER_TILE * 16 + ii * 16);
auto &B_sc_sm_subtile_2 = *reinterpret_cast<st_fp8e4m3<32, 16, false> *>(reinterpret_cast<uint64_t>(&input_scales[stage].sfB_2.data[0]) + 16 * 32 * ii);
load_mxnv_scale_async(B_sc_tm_subtile_2, B_sc_sm_subtile_2);
}
// asm volatile(
// "{tcgen05.commit.cta_group::1.mbarrier::arrive::one.shared::cluster.b64 [%0];}"
// :: "l"(__cvta_generic_to_shared(&inputs_all_arrived[stage]))); // todo: isn't there a wrapper?
kittens::detail::tcgen05::commit<1>(inputs_all_arrived[stage]);
stage = (stage + 1) % globals::PIPELINE_STAGES;
}
}
} else if (warp_id == 2 && lane_id == 0) {
for (int block_idx = cluster_id; block_idx < num_blocks; block_idx += gridDim.x / config::CLUSTER_SIZE) {
wait(tensor_finished, get_phasebit<1>(phasebits_2, globals::PIPELINE_STAGES));
update_phasebit<1>(phasebits_2, globals::PIPELINE_STAGES);
for (int i = 0; i < num_iters_per_block; i++) {
wait(inputs_all_arrived[stage], get_phasebit<0>(phasebits_2, stage));
update_phasebit<0>(phasebits_2, stage);
if (i == 0) mm_ABt(out_tm_2, input_tiles[stage].A, input_tiles[stage].B_2,
A_sc_tm.subtile<full_tt_fp8e4m3<globals::PACKED_PER_TILE * 16>>(stage * globals::PACKED_PER_TILE * 16),
B_sc_tm_2.subtile<full_tt_fp8e4m3<globals::PACKED_PER_TILE * 16>>(stage * globals::PACKED_PER_TILE * 16),
matmul_finished_2[stage]);
else mma_ABt(out_tm_2, input_tiles[stage].A, input_tiles[stage].B_2,
A_sc_tm.subtile<full_tt_fp8e4m3<globals::PACKED_PER_TILE * 16>>(stage * globals::PACKED_PER_TILE * 16),
B_sc_tm_2.subtile<full_tt_fp8e4m3<globals::PACKED_PER_TILE * 16>>(stage * globals::PACKED_PER_TILE * 16),
matmul_finished_2[stage]);
if (i == num_iters_per_block - 1) {
last_stage = stage;
}
stage = (stage + 1) % globals::PIPELINE_STAGES;
}
}
} else if (warp_id == 3 && lane_id == 0) {
for (int block_idx = cluster_id; block_idx < num_blocks; block_idx += gridDim.x / config::CLUSTER_SIZE) {
wait(tensor_finished, get_phasebit<1>(phasebits, globals::PIPELINE_STAGES));
update_phasebit<1>(phasebits, globals::PIPELINE_STAGES);
for (int i = 0; i < num_iters_per_block; i++) {
wait(inputs_all_arrived[stage], get_phasebit<0>(phasebits, stage));
update_phasebit<0>(phasebits, stage);
if (i == 0) mm_ABt(out_tm, input_tiles[stage].A, input_tiles[stage].B,
A_sc_tm.subtile<full_tt_fp8e4m3<globals::PACKED_PER_TILE * 16>>(stage * globals::PACKED_PER_TILE * 16),
B_sc_tm.subtile<full_tt_fp8e4m3<globals::PACKED_PER_TILE * 16>>(stage * globals::PACKED_PER_TILE * 16),
matmul_finished_1[stage]);
else mma_ABt(out_tm, input_tiles[stage].A, input_tiles[stage].B,
A_sc_tm.subtile<full_tt_fp8e4m3<globals::PACKED_PER_TILE * 16>>(stage * globals::PACKED_PER_TILE * 16),
B_sc_tm.subtile<full_tt_fp8e4m3<globals::PACKED_PER_TILE * 16>>(stage * globals::PACKED_PER_TILE * 16),
matmul_finished_1[stage]);
stage = (stage + 1) % globals::PIPELINE_STAGES;
}
}
}
}
else {
// warpgroup::increase_registers<256>();
for (int block_idx = cluster_id; block_idx < num_blocks;
block_idx += gridDim.x / config::CLUSTER_SIZE) {
int row_block_idx = block_idx / num_blocks_per_col;
int col_block_idx = block_idx % num_blocks_per_col;
wait(outputs_arrived,
get_phasebit<0>(phasebits, globals::PIPELINE_STAGES));
// wait(outputs_arrived, get_phasebit<0>(phasebits_2,
// globals::PIPELINE_STAGES));
update_phasebit<0>(phasebits, globals::PIPELINE_STAGES);
update_phasebit<0>(phasebits_2, globals::PIPELINE_STAGES);
rt_fl<globals::TILE_SIZE / 4, globals::TILE_SIZE> C_reg_2;
rt_fl<globals::TILE_SIZE / 4, globals::TILE_SIZE> C_reg_1;
warpgroup::load_async(C_reg_1, out_tm);
tensor_load_wait(); // are those syncs necessary?
warpgroup::sync(1);
warpgroup::load_async(C_reg_2, out_tm_2);
tensor_load_wait();
warpgroup::sync(1);
// todo: explicit unrolling to guide compiler
#pragma unroll // TODO: software emulate exp on CUDA Cores
for (int i = 0; i < C_reg_1.height; i++) {
#pragma unroll
for (int j = 0; j < C_reg_1.width; j++) {
#pragma unroll
for (int k = 0; k < C_reg_1.packed_per_tile; k++) {
float x = C_reg_1.tiles[i][j].data[k].x;
float y = C_reg_1.tiles[i][j].data[k].y;
float sx = (x >= 0.0f)
? 1.0f / (1.0f + expf(-x))
: expf(x) / (1.0f + expf(x)); // todo: packed instructions?
float sy = (y >= 0.0f)
? 1.0f / (1.0f + expf(-y))
: expf(y) / (1.0f + expf(y));
float silu_x = x * sx;
float silu_y = y * sy;
float gate_x = C_reg_2.tiles[i][j].data[k].x;
float gate_y = C_reg_2.tiles[i][j].data[k].y;
C_reg_1.tiles[i][j].data[k].x = silu_x * gate_x;
C_reg_1.tiles[i][j].data[k].y = silu_y * gate_y;
}
}
}
if (warpgroup::laneid() == 0) arrive(tensor_finished, 1); // todo: can this be done earlier?
warpgroup::store(output_tiles.C, C_reg_1);
warpgroup::sync(1);
warpgroup::tma::store_async(G.C, output_tiles.C, {row_block_idx, col_block_idx});
warpgroup::tma::store_async_read_wait();
}
}
}
void entrypoint(const at::Tensor &A, const at::Tensor &B,
const at::Tensor &B_2, at::Tensor &C,
const at::Tensor &sfA, const at::Tensor &sfB,
const at::Tensor &sfB_2) {
auto sfA_swizzled = sfA.view({sfA.size(0) / 128, 128, sfA.size(1) / 4, 4})
.transpose(1,2)
.reshape({sfA.size(0)/128, sfA.size(1)/4, 4, 32, 4})
.transpose(-2,-3)
.reshape({sfA.size(0)/128, sfA.size(1)/4, 32, 16})
.reshape({sfA.size(0)/128, (sfA.size(1)/4)/4, 128, 16});
auto sfB_swizzled = sfB.view({sfB.size(0) / 128, 128, sfB.size(1) / 4, 4})
.transpose(1,2)
.reshape({sfB.size(0)/128, sfB.size(1)/4, 4, 32, 4})
.transpose(-2,-3)
.reshape({sfB.size(0)/128, sfB.size(1)/4, 32, 16})
.reshape({sfB.size(0)/128, (sfB.size(1)/4)/4, 128, 16});
auto sfB2_swizzled = sfB_2.view({sfB_2.size(0) / 128, 128, sfB_2.size(1) / 4, 4})
.transpose(1,2)
.reshape({sfB_2.size(0)/128, sfB_2.size(1)/4, 4, 32, 4})
.transpose(-2,-3)
.reshape({sfB_2.size(0)/128, sfB_2.size(1)/4, 32, 16})
.reshape({sfB_2.size(0)/128, (sfB_2.size(1)/4)/4, 128, 16});
globals G{
.A = kittens::py::tensor_to_gl<globals::A_fp4x2_gl>(A),
.B = kittens::py::tensor_to_gl<globals::B_fp4x2_gl>(B),
.B_2 = kittens::py::tensor_to_gl<globals::B_fp4x2_gl>(B_2),
.C = kittens::py::tensor_to_gl<globals::C_gl>(C),
.sfA = kittens::py::tensor_to_gl<globals::sfA_gl>(sfA_swizzled),
.sfB = kittens::py::tensor_to_gl<globals::sfB_gl>(sfB_swizzled),
.sfB_2 = kittens::py::tensor_to_gl<globals::sfB_gl>(sfB2_swizzled)
};
kittens::py::launch_kernel<config, globals, kernel>(G);
}
} // namespace nvfp4_gemm
PYBIND11_MODULE(_C, m) {
m.doc() = "nvpf4 python module";
m.def("kernel", &nvfp4_gemm::entrypoint);
}
'''
include_flags = [
f'-I{thunderkittens_dir}/include',
f'-I{thunderkittens_dir}/prototype',
]
extra_cflags = ['-O3','-std=c++20','-DNDEBUG','-fPIE','-fopenmp','-Wno-psabi','-fno-strict-aliasing','-fPIC'] + include_flags
extra_cuda_cflags = [
'-O3','-std=c++20','-DNDEBUG',
'--expt-extended-lambda','--expt-relaxed-constexpr',
'-Xcompiler=-fPIE','-Xcompiler=-fopenmp',
'-Xcompiler=-Wno-psabi','-Xcompiler=-fno-strict-aliasing',
'--use_fast_math','-forward-unknown-to-host-compiler',
'-Xcompiler=-fPIC',
# '-DKITTENS_HOPPER',
'-DKITTENS_BLACKWELL',
'-D__CUDA_NO_HALF_OPERATORS__',
'-D__CUDA_NO_HALF_CONVERSIONS__',
'-D__CUDA_NO_BFLOAT16_CONVERSIONS__',
'-D__CUDA_NO_HALF2_OPERATORS__',
'-gencode=arch=compute_100a,code=sm_100a'
] + include_flags
extra_ldflags = [
'-L/usr/local/cuda/lib64',
'-lcuda',
'-lcudart',
'-L/home/runner/competition/cuda/lib/python3.12/site-packages/torch/lib',
'-lc10', '-lc10_cuda', '-ltorch_cpu', '-ltorch_cuda', '-ltorch', '-ltorch_python',
'-Wl,-rpath,/usr/local/cuda/lib64',
]
nvpf4 = load_inline(
name='_C',
cpp_sources='',
cuda_sources=nvpf4_src,
extra_cflags=extra_cflags,
extra_cuda_cflags=extra_cuda_cflags,
extra_ldflags=extra_ldflags,
verbose=True,
with_cuda=True,
)
def scale_swizzle(
V_sc_unswizzled: torch.Tensor, # (M, N // 16) fp8e4m3
packed_per_tile: int # number of packed K=64 blocks per scale tile
) -> torch.Tensor:
# https://docs.nvidia.com/cuda/parallel-thread-execution/#tcgen05-mma-scale-factor-a-layout-4x
assert len(V_sc_unswizzled.shape) == 2
assert V_sc_unswizzled.dtype == torch.float8_e4m3fn
assert V_sc_unswizzled.shape[0] % 128 == 0
assert (V_sc_unswizzled.shape[1] * 16) % 64 == 0
M_BLOCK = 128
N_BLOCK = 4 # 64 / 16
M, N_16 = V_sc_unswizzled.shape
V_sc = V_sc_unswizzled # (M, N_16)
V_sc = V_sc.reshape( # (M / 128, 128, N_16 / 4, 4)
M // M_BLOCK, M_BLOCK,
N_16 // N_BLOCK, N_BLOCK
)
V_sc = V_sc.transpose(1, 2) # (M / 128, N_16 / 4, 128, 4) --> last 2 dims are all we need per MM
V_sc = V_sc.reshape( # (M / 128, N_16 / 4, 4, 32, 4)
M // M_BLOCK, N_16 // N_BLOCK,
4, M_BLOCK // 4, N_BLOCK
)
V_sc = V_sc.transpose(-2, -3) # (M / 128, N_16 / 4, 32, 4, 4)
V_sc = V_sc.reshape( # (M / 128, N_16 / 4, 32, 16)
M // M_BLOCK,
N_16 // N_BLOCK,
M_BLOCK // 4, N_BLOCK * 4
)
V_sc = V_sc.reshape( # (M / 128, N_16 / 4, 32, 16)
M // M_BLOCK, # Pack the scale tiles. This is purely for GEMM efficiency
N_16 // N_BLOCK // packed_per_tile,
packed_per_tile * M_BLOCK // 4, N_BLOCK * 4
)
return V_sc
def ceil_div(a, b):
return (a + b - 1) // b
def to_blocked(input_matrix):
rows, cols = input_matrix.shape
n_row_blocks = ceil_div(rows, 128)
n_col_blocks = ceil_div(cols, 4)
padded = input_matrix
blocks = padded.view(n_row_blocks, 128, n_col_blocks, 4).permute(0, 2, 1, 3)
rearranged = blocks.reshape(-1, 4, 32, 4).transpose(1, 2).reshape(-1, 32, 16)
return rearranged.flatten()
def ref_kernel(
data: input_t,
) -> output_t:
"""
PyTorch reference implementation of NVFP4 block-scaled dual GEMM with silu activation,
C = silu(A @ B1) * (A @ B2).
"""
a_ref, b1_ref, b2_ref, sfa_ref_cpu, sfb1_ref_cpu, sfb2_ref_cpu, _, _, _, c_ref = data
# Get dimensions from MxNxL layout
m, n, l = c_ref.shape
# Call torch._scaled_mm to compute the GEMV result
ref1 = torch.empty(
(l, m, n),
dtype=torch.float32,
device="cuda",
).permute(1, 2, 0)
ref2 = torch.empty(
(l, m, n),
dtype=torch.float32,
device="cuda",
).permute(1, 2, 0)
for l_idx in range(l):
# Convert the scale factor tensor to blocked format
scale_a = to_blocked(sfa_ref_cpu[:, :, l_idx])
scale_b1 = to_blocked(sfb1_ref_cpu[:, :, l_idx])
scale_b2 = to_blocked(sfb2_ref_cpu[:, :, l_idx])
# (m, k) @ (n, k).T -> (m, n)
res1 = torch._scaled_mm(
a_ref[:, :, l_idx],
b1_ref[:, :, l_idx].transpose(0, 1),
scale_a.cuda(),
scale_b1.cuda(),
# torch.ones_like(scale_a, device="cuda"),
# torch.ones_like(scale_b1, device="cuda"),
bias=None,
out_dtype=torch.float32,
)
ref1[:, :, l_idx] = res1
res2 = torch._scaled_mm(
a_ref[:, :, l_idx],
b2_ref[:, :, l_idx].transpose(0, 1),
scale_a.cuda(),
scale_b2.cuda(),
# torch.ones_like(scale_a, device="cuda"),
# torch.ones_like(scale_b2, device="cuda"),
bias=None,
out_dtype=torch.float32,
)
ref2[:, :, l_idx] = res2
# Do silu on the first GEMM result and multiply with the second GEMM result
c_ref = (torch.nn.functional.silu(ref1) * ref2).to(torch.float16)
# c_ref = ref1
return c_ref
def custom_kernel(data: input_t) -> output_t:
a_ref, b1_ref, b2_ref, sfa_ref_cpu, sfb1_ref_cpu, sfb2_ref_cpu, sfa_cute, sfb1_cute, sfb2_cute, c_ref = data
if a_ref.shape[1] < 256:
out = ref_kernel(data)
return out
else:
nvpf4.kernel(a_ref, b1_ref, b2_ref, c_ref, sfa_ref_cpu, sfb1_ref_cpu, sfb2_ref_cpu)
return c_ref
scrolls · 574 lines total
Source code from GPU Mode and the KernelBot dataset · June 9 Researcher Reciprocity License v1.0
Best evidence level for this revision: reported
JSON