submission 75298
sohail · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 1573 lines, June 9 Researcher Reciprocity License v1.0.
submission_tc_tmem.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-nvfp4-gemv-75298?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:a5086b6b664f3e20f27d0425786fba41c21139f9003c7b7321f4a5f0df6fca78
license declaredunknown
license concludedunknown
authorssohail
imported2026-08-26
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
async-copy
__device__ __forceinline__ uint32_t tcgen_cp_async(fp4
printf("[nvfp4] selected kernel path = %d\n", static_cast<int>(selected));mbarrier
__device__ __forceinline__ void mbarrier_init_cta(uint64_t* bar, uint32_t pending_count) {mma
namespace wmma = nvcuda::wmma;shared-memory
__device__ __forceinline__ uint32_t smem_ptr(const void* ptr) {tcgen05
"tcgen05.alloc.cta_group::1.sync.aligned.shared::cta.b32 [%0], %1;\n"Kernel source
submission_tc_tmem.py1573 lines
# submission.py
# Optimized PyTorch extension for B200 NVFP4 operations with FP8 scales
import os
import torch
from torch.utils.cpp_extension import load_inline
_EXT = None # lazily compiled extension
_FP4_LUT_CACHE = {}
def _get_fp4_lut(device):
key = device.index if device.type == "cuda" else device
lut = _FP4_LUT_CACHE.get(key)
if lut is None:
values = torch.tensor(
[
0.0, 0.5, 1.0, 1.5, 2.0, 3.0, 4.0, 6.0,
0.0, -0.5, -1.0, -1.5, -2.0, -3.0, -4.0, -6.0,
],
device=device,
dtype=torch.float32,
)
_FP4_LUT_CACHE[key] = values
lut = values
return lut
def _decode_fp4_tensor(packed):
lut = _get_fp4_lut(packed.device)
low = (packed & 0xF).to(torch.long)
high = ((packed >> 4) & 0xF).to(torch.long)
decoded = torch.stack((lut[low], lut[high]), dim=-1)
shape = packed.shape
return decoded.view(shape[0], shape[1] * 2, shape[2]).contiguous()
def _decode_fp8_tensor(packed):
if hasattr(torch, "float8_e4m3fn"):
return packed.view(torch.float8_e4m3fn).to(torch.float32).contiguous()
# Fallback decoding if float8 dtype is unavailable
bytes_float = packed.to(torch.float32)
sign = ((packed >> 7) & 0x1).to(torch.float32)
exponent = ((packed >> 3) & 0xF).to(torch.float32)
mantissa = (packed & 0x7).to(torch.float32)
is_zero = (exponent == 0) & (mantissa == 0)
norm = exponent >= 1
value = torch.empty_like(bytes_float)
value[norm] = (1.0 + mantissa[norm] / 8.0) * torch.pow(2.0, exponent[norm] - 7.0)
value[~norm] = (mantissa[~norm] / 8.0) * torch.pow(2.0, -6.0)
value[is_zero] = 0.0
return torch.where(sign > 0, -value, value).contiguous()
def _get_ext():
global _EXT
if _EXT is not None:
return _EXT
# Minimal C++ source with forward declaration
CPP_SRC = """
#include <torch/extension.h>
void bmv_launcher(
torch::Tensor A,
torch::Tensor B,
torch::Tensor SFA,
torch::Tensor SFB,
torch::Tensor C);
"""
enable_tcgen_build = os.environ.get("NVFP4_ENABLE_TCGEN_BUILD", "").lower() not in {"", "0", "false", "no"}
experimental_tcgen = os.environ.get("NVFP4_EXPERIMENTAL_TCGEN_MMA", "").lower() not in {"", "0", "false", "no"}
if experimental_tcgen:
print("Enabling NVFP4_EXPERIMENTAL_TCGEN_MMA build flag (1)")
if os.environ.get("NVFP4_CAPTURE_TMEM_DEBUG"):
keep_flags = "-keep --keep-dir /tmp/nvcc_keep"
existing = os.environ.get("TORCH_NVCC_FLAGS", "")
if keep_flags not in existing:
combined = f"{existing} {keep_flags}".strip()
os.environ["TORCH_NVCC_FLAGS"] = combined
os.makedirs("/tmp/nvcc_keep", exist_ok=True)
# CUDA implementation
CUDA_SRC = r"""
#include <cuda.h>
#include <cuda_runtime.h>
#include <cuda_fp16.h>
#include <mma.h>
#include <torch/extension.h>
#include <ATen/cuda/CUDAContext.h>
#include <c10/cuda/CUDAStream.h>
#include <stdint.h>
#include <cstdlib>
#include <cstring>
#include <cctype>
#include <cstdio>
#include <cmath>
#include <vector>
#ifndef NVFP4_ENABLE_TCGEN
#define NVFP4_ENABLE_TCGEN 0
#endif
#ifndef NVFP4_EXPERIMENTAL_TCGEN_MMA
#define NVFP4_EXPERIMENTAL_TCGEN_MMA 0
#endif
#ifndef ENABLE_WMMA_KERNEL
#define ENABLE_WMMA_KERNEL 0
#endif
#if defined(__CUDA_ARCH__) && __CUDA_ARCH__ >= 800
#define CP_ASYNC_SUPPORTED 1
#else
#define CP_ASYNC_SUPPORTED 0
#endif
namespace wmma = nvcuda::wmma;
constexpr int WMMA_M = 16;
constexpr int WMMA_N = 8;
constexpr int WMMA_K = 64;
constexpr int CTA_M = WMMA_M * 8;
constexpr int CTA_N = WMMA_N;
constexpr int CTA_WARPS = CTA_M / WMMA_M;
constexpr int TCGEN_CTA_M = WMMA_M * 4; // 64 rows per CTA (one warpgroup)
constexpr int TCGEN_CTA_WARPS = TCGEN_CTA_M / WMMA_M;
constexpr int TCGEN_THREADS = TCGEN_CTA_WARPS * 32;
enum class KernelPath : int {
Auto = 0,
Scalar = 1,
WMMA = 2,
TCGEN = 3
};
enum class SmallNMode : int {
Auto = 0,
Batch = 1,
Pad = 2,
Scalar = 3
};
inline bool equals_ignore_case(const char* lhs, const char* rhs) {
if (lhs == nullptr || rhs == nullptr) {
return false;
}
while (*lhs && *rhs) {
if (std::tolower(*lhs) != std::tolower(*rhs)) {
return false;
}
++lhs;
++rhs;
}
return (*lhs == '\0') && (*rhs == '\0');
}
inline KernelPath kernel_path_override_from_env() {
const char* env = std::getenv("NVFP4_FORCE_KERNEL");
if (env == nullptr || env[0] == '\0') {
return KernelPath::Auto;
}
if (equals_ignore_case(env, "scalar")) {
return KernelPath::Scalar;
}
if (equals_ignore_case(env, "wmma")) {
return KernelPath::WMMA;
}
if (equals_ignore_case(env, "tcgen") || equals_ignore_case(env, "tmemory")) {
return KernelPath::TCGEN;
}
return KernelPath::Auto;
}
inline SmallNMode small_n_mode_override_from_env() {
const char* env = std::getenv("NVFP4_SMALLN_MODE");
if (env == nullptr || env[0] == '\0') {
return SmallNMode::Auto;
}
if (equals_ignore_case(env, "batch")) {
return SmallNMode::Batch;
}
if (equals_ignore_case(env, "pad")) {
return SmallNMode::Pad;
}
if (equals_ignore_case(env, "scalar")) {
return SmallNMode::Scalar;
}
return SmallNMode::Auto;
}
inline bool tcgen_auto_enabled() {
const char* env = std::getenv("NVFP4_ENABLE_TCGEN_AUTO");
return env && env[0] != '\0' && env[0] != '0';
}
#if NVFP4_EXPERIMENTAL_TCGEN_MMA
inline bool capture_tmem_debug_enabled() {
const char* env = std::getenv("NVFP4_CAPTURE_TMEM_DEBUG");
return env && env[0] != '\0' && env[0] != '0';
}
#endif
inline bool arch_supports_wmma(const cudaDeviceProp* prop) {
return ENABLE_WMMA_KERNEL && prop && prop->major >= 8;
}
inline bool arch_supports_tcgen(const cudaDeviceProp* prop) {
return prop && prop->major >= 10;
}
// -----------------------------------------------
// Tensor Memory helpers (alloc, dealloc, copies)
// -----------------------------------------------
__device__ __forceinline__ uint32_t smem_ptr(const void* ptr) {
return static_cast<uint32_t>(__cvta_generic_to_shared(ptr));
}
#if NVFP4_ENABLE_TCGEN
__device__ __forceinline__ void tcgen_alloc_cols(uint32_t* dst_smem_ptr, uint32_t num_cols) {
asm volatile(
"tcgen05.alloc.cta_group::1.sync.aligned.shared::cta.b32 [%0], %1;\n"
:
: "r"(smem_ptr(dst_smem_ptr)), "r"(num_cols)
: "memory"
);
}
__device__ __forceinline__ void tcgen_dealloc_cols(uint32_t tmem_addr, uint32_t num_cols) {
asm volatile(
"tcgen05.dealloc.cta_group::1.sync.aligned.b32 %0, %1;\n"
:
: "r"(tmem_addr), "r"(num_cols)
: "memory"
);
}
#else
__device__ __forceinline__ void tcgen_alloc_cols(uint32_t*, uint32_t) {}
__device__ __forceinline__ void tcgen_dealloc_cols(uint32_t, uint32_t) {}
#endif
__device__ __forceinline__ uint64_t encode_matrix_component(uint32_t value_bytes) {
return static_cast<uint64_t>((value_bytes & 0x3FFFFu) >> 4);
}
__device__ __forceinline__ uint64_t make_smem_matrix_desc(
uint32_t smem_addr_bytes,
uint32_t lead_dim_bytes,
uint32_t stride_dim_bytes,
bool lead_is_absolute,
int swizzle_mode = 0)
{
uint64_t desc = 0;
desc |= encode_matrix_component(smem_addr_bytes) << 0;
desc |= encode_matrix_component(lead_dim_bytes) << 16;
desc |= encode_matrix_component(stride_dim_bytes) << 32;
desc |= (uint64_t)0x1 << 46; // fixed constant 0b001
desc |= (uint64_t)(lead_is_absolute ? 1 : 0) << 52;
desc |= (uint64_t)0xB0 << 53; // fixed constant per spec
desc |= (uint64_t)(swizzle_mode & 0x7) << 61;
return desc;
}
#if NVFP4_ENABLE_TCGEN
__device__ __forceinline__ uint32_t tcgen_cp_async(
uint32_t tmem_addr,
uint64_t smem_desc,
int shape_selector)
{
switch (shape_selector) {
case 0:
asm volatile(
"tcgen05.cp.cta_group::1.128x256b.b8x16.b4x16_p64 [%0], %1;\n"
:
: "r"(tmem_addr), "l"(smem_desc)
: "memory"
);
break;
default:
break;
}
return tmem_addr;
}
#else
__device__ __forceinline__ uint32_t tcgen_cp_async(uint32_t tmem_addr, uint64_t, int) {
return tmem_addr;
}
#endif
#if NVFP4_ENABLE_TCGEN
__device__ __forceinline__ uint32_t tmem_add_lane_offset(uint32_t base_addr, uint32_t lane_offset) {
const uint32_t column = base_addr & 0xFFFFu;
const uint32_t lane = ((base_addr >> 16) + lane_offset) & 0xFFFFu;
return (lane << 16) | column;
}
#if NVFP4_EXPERIMENTAL_TCGEN_MMA
__device__ float g_debug_tm_frag[TCGEN_THREADS * 4];
__device__ float g_debug_ref_frag[TCGEN_THREADS * 4];
#endif
#else
__device__ __forceinline__ uint32_t tmem_add_lane_offset(uint32_t base_addr, uint32_t) {
return base_addr;
}
#endif
__device__ __forceinline__ void mbarrier_init_cta(uint64_t* bar, uint32_t pending_count) {
#if defined(__CUDA_ARCH__)
asm volatile(
"mbarrier.init.shared::cta.b64 [%0], %1;\n"
:
: "r"(smem_ptr(bar)), "r"(pending_count)
: "memory"
);
#else
(void)bar;
(void)pending_count;
#endif
}
__device__ __forceinline__ void tcgen_commit_mbarrier(uint64_t* bar) {
#if defined(__CUDA_ARCH__) && defined(__CUDA_ARCH_FAMILY_SPECIFIC__)
asm volatile(
"tcgen05.commit.cta_group::1.mbarrier::arrive::one.b64 [%0];\n"
:
: "r"(smem_ptr(bar))
: "memory"
);
#else
(void)bar;
#endif
}
__device__ __forceinline__ void mbarrier_wait_parity_cta(uint64_t* bar, uint32_t parity) {
#if defined(__CUDA_ARCH__)
unsigned int complete = 0;
do {
asm volatile(
"{\n\t"
".reg .pred p;\n\t"
"mbarrier.try_wait.parity.shared::cta.b64 p, [%1], %2;\n\t"
"selp.u32 %0, 1, 0, p;\n\t"
"}\n"
: "=r"(complete)
: "r"(smem_ptr(bar)), "r"(parity)
: "memory"
);
} while (!complete);
#else
(void)bar;
(void)parity;
#endif
}
__device__ __forceinline__ uint32_t make_mxf4nvf4_instr_desc(
int m_dim,
int n_dim,
bool scale_is_ue4m3,
int scaleA_id,
int scaleB_id)
{
uint32_t desc = 0;
desc |= (scaleB_id & 0x3) << 4;
desc |= (1u) << 7; // atype = E2M1
desc |= (1u) << 10; // btype = E2M1
desc |= (uint32_t)((n_dim >> 3) & 0x3F) << 17;
desc |= (scale_is_ue4m3 ? 0u : 1u) << 23;
desc |= (uint32_t)((m_dim >> 7) & 0x3) << 27;
desc |= (scaleA_id & 0x3) << 29;
return desc;
}
inline bool tcgen_small_n_allowed(int L, SmallNMode mode) {
switch (mode) {
case SmallNMode::Batch:
return L >= WMMA_N;
case SmallNMode::Pad:
return L >= 1;
case SmallNMode::Scalar:
return false;
case SmallNMode::Auto:
default:
return L >= 4;
}
}
inline KernelPath choose_kernel_path(
const cudaDeviceProp* prop,
KernelPath override,
SmallNMode small_n_mode,
int K_bytes,
int L)
{
if (override != KernelPath::Auto) {
return override;
}
const bool k_aligned = ((K_bytes * 2) % WMMA_K) == 0;
if (tcgen_auto_enabled() && arch_supports_tcgen(prop) && k_aligned && tcgen_small_n_allowed(L, small_n_mode)) {
#if NVFP4_ENABLE_TCGEN
return KernelPath::TCGEN;
#endif
}
if (arch_supports_wmma(prop) && (K_bytes % (WMMA_K / 2) == 0) && L >= CTA_N) {
return KernelPath::WMMA;
}
return KernelPath::Scalar;
}
// ----------------------- FP4/FP8 converters -----------------------
__device__ __constant__ float FP4_E2M1_LUT[16] = {
0.0f, 0.5f, 1.0f, 1.5f, 2.0f, 3.0f, 4.0f, 6.0f,
0.0f, -0.5f, -1.0f, -1.5f, -2.0f, -3.0f, -4.0f, -6.0f
};
__device__ __constant__ float FP8_E4M3_LUT[256] = {
0.0f, 0.0019531250f, 0.0039062500f, 0.0058593750f,
0.0078125000f, 0.0097656250f, 0.0117187500f, 0.0136718750f,
0.0156250000f, 0.0175781250f, 0.0195312500f, 0.0214843750f,
0.0234375000f, 0.0253906250f, 0.0273437500f, 0.0292968750f,
0.0312500000f, 0.0351562500f, 0.0390625000f, 0.0429687500f,
0.0468750000f, 0.0507812500f, 0.0546875000f, 0.0585937500f,
0.0625000000f, 0.0703125000f, 0.0781250000f, 0.0859375000f,
0.0937500000f, 0.1015625000f, 0.1093750000f, 0.1171875000f,
0.1250000000f, 0.1406250000f, 0.1562500000f, 0.1718750000f,
0.1875000000f, 0.2031250000f, 0.2187500000f, 0.2343750000f,
0.2500000000f, 0.2812500000f, 0.3125000000f, 0.3437500000f,
0.3750000000f, 0.4062500000f, 0.4375000000f, 0.4687500000f,
0.5000000000f, 0.5625000000f, 0.6250000000f, 0.6875000000f,
0.7500000000f, 0.8125000000f, 0.8750000000f, 0.9375000000f,
1.0000000000f, 1.1250000000f, 1.2500000000f, 1.3750000000f,
1.5000000000f, 1.6250000000f, 1.7500000000f, 1.8750000000f,
2.0000000000f, 2.2500000000f, 2.5000000000f, 2.7500000000f,
3.0000000000f, 3.2500000000f, 3.5000000000f, 3.7500000000f,
4.0000000000f, 4.5000000000f, 5.0000000000f, 5.5000000000f,
6.0000000000f, 6.5000000000f, 7.0000000000f, 7.5000000000f,
8.0000000000f, 9.0000000000f, 10.0000000000f, 11.0000000000f,
12.0000000000f, 13.0000000000f, 14.0000000000f, 15.0000000000f,
16.0000000000f, 18.0000000000f, 20.0000000000f, 22.0000000000f,
24.0000000000f, 26.0000000000f, 28.0000000000f, 30.0000000000f,
32.0000000000f, 36.0000000000f, 40.0000000000f, 44.0000000000f,
48.0000000000f, 52.0000000000f, 56.0000000000f, 60.0000000000f,
64.0000000000f, 72.0000000000f, 80.0000000000f, 88.0000000000f,
96.0000000000f, 104.0000000000f, 112.0000000000f, 120.0000000000f,
128.0000000000f, 144.0000000000f, 160.0000000000f, 176.0000000000f,
192.0000000000f, 208.0000000000f, 224.0000000000f, 240.0000000000f,
256.0000000000f, 288.0000000000f, 320.0000000000f, 352.0000000000f,
384.0000000000f, 416.0000000000f, 448.0000000000f, 448.0000000000f,
-0.0f, -0.0019531250f, -0.0039062500f, -0.0058593750f,
-0.0078125000f, -0.0097656250f, -0.0117187500f, -0.0136718750f,
-0.0156250000f, -0.0175781250f, -0.0195312500f, -0.0214843750f,
-0.0234375000f, -0.0253906250f, -0.0273437500f, -0.0292968750f,
-0.0312500000f, -0.0351562500f, -0.0390625000f, -0.0429687500f,
-0.0468750000f, -0.0507812500f, -0.0546875000f, -0.0585937500f,
-0.0625000000f, -0.0703125000f, -0.0781250000f, -0.0859375000f,
-0.0937500000f, -0.1015625000f, -0.1093750000f, -0.1171875000f,
-0.1250000000f, -0.1406250000f, -0.1562500000f, -0.1718750000f,
-0.1875000000f, -0.2031250000f, -0.2187500000f, -0.2343750000f,
-0.2500000000f, -0.2812500000f, -0.3125000000f, -0.3437500000f,
-0.3750000000f, -0.4062500000f, -0.4375000000f, -0.4687500000f,
-0.5000000000f, -0.5625000000f, -0.6250000000f, -0.6875000000f,
-0.7500000000f, -0.8125000000f, -0.8750000000f, -0.9375000000f,
-1.0000000000f, -1.1250000000f, -1.2500000000f, -1.3750000000f,
-1.5000000000f, -1.6250000000f, -1.7500000000f, -1.8750000000f,
-2.0000000000f, -2.2500000000f, -2.5000000000f, -2.7500000000f,
-3.0000000000f, -3.2500000000f, -3.5000000000f, -3.7500000000f,
-4.0000000000f, -4.5000000000f, -5.0000000000f, -5.5000000000f,
-6.0000000000f, -6.5000000000f, -7.0000000000f, -7.5000000000f,
-8.0000000000f, -9.0000000000f, -10.0000000000f, -11.0000000000f,
-12.0000000000f, -13.0000000000f, -14.0000000000f, -15.0000000000f,
-16.0000000000f, -18.0000000000f, -20.0000000000f, -22.0000000000f,
-24.0000000000f, -26.0000000000f, -28.0000000000f, -30.0000000000f,
-32.0000000000f, -36.0000000000f, -40.0000000000f, -44.0000000000f,
-48.0000000000f, -52.0000000000f, -56.0000000000f, -60.0000000000f,
-64.0000000000f, -72.0000000000f, -80.0000000000f, -88.0000000000f,
-96.0000000000f, -104.0000000000f, -112.0000000000f, -120.0000000000f,
-128.0000000000f, -144.0000000000f, -160.0000000000f, -176.0000000000f,
-192.0000000000f, -208.0000000000f, -224.0000000000f, -240.0000000000f,
-256.0000000000f, -288.0000000000f, -320.0000000000f, -352.0000000000f,
-384.0000000000f, -416.0000000000f, -448.0000000000f, -448.0000000000f
};
__device__ __forceinline__ float decode_fp4_e2m1(uint8_t nibble) {
return FP4_E2M1_LUT[nibble & 0xF];
}
__device__ __forceinline__ float decode_fp8_e4m3(uint8_t byte_val) {
return FP8_E4M3_LUT[byte_val];
}
__device__ __forceinline__ int frag_a_row(int group_id, int elem_idx) {
return ((elem_idx < 8) || (elem_idx >= 16 && elem_idx < 24)) ? group_id : (group_id + 8);
}
__device__ __forceinline__ int frag_a_col(int tid_in_group, int elem_idx) {
int col = tid_in_group * 8 + (elem_idx & 7);
if (elem_idx >= 16) {
col += 32;
}
return col;
}
__device__ __forceinline__ int frag_b_row(int tid_in_group, int elem_idx) {
int row = tid_in_group * 8 + (elem_idx & 7);
if (elem_idx >= 8) {
row += 32;
}
return row;
}
__device__ __forceinline__ int frag_b_col(int group_id) {
return group_id;
}
__device__ __forceinline__ int frag_acc_row(int group_id, int acc_idx) {
return (acc_idx < 2) ? group_id : (group_id + 8);
}
__device__ __forceinline__ int frag_acc_col(int tid_in_group, int acc_idx) {
return tid_in_group * 2 + (acc_idx & 1);
}
__device__ __forceinline__ uint8_t load_stage_a_nibble(
const uint8_t* tile,
int row_local,
int col_local)
{
const int bytes_per_row = WMMA_K / 2;
const uint8_t byte_val = tile[row_local * bytes_per_row + (col_local >> 1)];
return (col_local & 1) ? ((byte_val >> 4) & 0xF) : (byte_val & 0xF);
}
__device__ __forceinline__ uint8_t load_stage_b_nibble(
const uint8_t* tile,
int col_local)
{
const uint8_t byte_val = tile[col_local >> 1];
return (col_local & 1) ? ((byte_val >> 4) & 0xF) : (byte_val & 0xF);
}
__device__ __forceinline__ uint8_t load_stage_sfa_byte(
const uint8_t* tile,
int row_local,
int elem)
{
return tile[row_local * 4 + elem];
}
__device__ __forceinline__ uint8_t load_stage_sfb_byte(
const uint8_t* tile,
int elem)
{
return tile[elem];
}
__device__ __forceinline__ uint8_t load_fp4_nibble(
const uint8_t* __restrict__ tensor,
int dim_m,
int dim_l,
int K_vals,
int global_row,
int global_k,
int global_l,
int64_t stride_m,
int64_t stride_k,
int64_t stride_l)
{
if (global_row < 0 || global_row >= dim_m || global_l < 0 || global_l >= dim_l) {
return 0;
}
if (global_k < 0 || global_k >= K_vals) {
return 0;
}
const int byte_index = global_k >> 1;
const bool high = (global_k & 1);
const int64_t offset = (int64_t)global_row * stride_m +
(int64_t)byte_index * stride_k +
(int64_t)global_l * stride_l;
const uint8_t byte_val = tensor[offset];
return high ? ((byte_val >> 4) & 0xF) : (byte_val & 0xF);
}
__device__ __forceinline__ uint8_t load_fp4_nibble_linear(
const uint8_t* data,
int K_vals,
int global_k)
{
if (global_k < 0 || global_k >= K_vals) {
return 0;
}
const int byte_index = global_k >> 1;
const bool high = (global_k & 1);
const uint8_t byte_val = data[byte_index];
return high ? ((byte_val >> 4) & 0xF) : (byte_val & 0xF);
}
__device__ __forceinline__ uint8_t load_scale_byte(
const uint8_t* __restrict__ tensor,
int dim_m,
int dim_l,
int K_scales,
int global_row,
int scale_idx,
int global_l,
int64_t stride_m,
int64_t stride_k,
int64_t stride_l)
{
if (global_row < 0 || global_row >= dim_m || global_l < 0 || global_l >= dim_l) {
return 0;
}
if (scale_idx < 0 || scale_idx >= K_scales) {
return 0;
}
const int64_t offset = (int64_t)global_row * stride_m +
(int64_t)scale_idx * stride_k +
(int64_t)global_l * stride_l;
return tensor[offset];
}
__global__ void bmv_kernel_scalar(
const uint8_t* __restrict__ A,
const uint8_t* __restrict__ B,
const uint8_t* __restrict__ SFA,
const uint8_t* __restrict__ SFB,
__half* __restrict__ C,
int M, int M_B, int K_bytes, int L,
int64_t a_stride_m, int64_t a_stride_k, int64_t a_stride_l,
int64_t b_stride_m, int64_t b_stride_k, int64_t b_stride_l,
int64_t sfa_stride_m, int64_t sfa_stride_k, int64_t sfa_stride_l,
int64_t sfb_stride_m, int64_t sfb_stride_k, int64_t sfb_stride_l)
{
const int l = blockIdx.y;
if (l >= L) {
return;
}
const int m = blockIdx.x * blockDim.x + threadIdx.x;
const int K_scales = K_bytes / 8;
if (K_scales == 0) {
return;
}
const int m_b = 0; // competition inputs broadcast B/SFB across M
extern __shared__ uint8_t shared_bytes[];
uint8_t* sh_B = shared_bytes;
const size_t sh_B_size = ((size_t)K_bytes + 15) & ~size_t(15);
float* sh_SFB = reinterpret_cast<float*>(shared_bytes + sh_B_size);
const int64_t b_row_base = (int64_t)m_b * b_stride_m + (int64_t)l * b_stride_l;
const int64_t sfb_row_base = (int64_t)m_b * sfb_stride_m + (int64_t)l * sfb_stride_l;
#if CP_ASYNC_SUPPORTED
const bool can_async = (b_stride_k == 1);
if (can_async) {
const int chunk = 16;
const int stride = blockDim.x * chunk;
for (int idx = threadIdx.x * chunk; idx + chunk <= K_bytes; idx += stride) {
void* dst = sh_B + idx;
const void* src = B + b_row_base + idx;
unsigned smem_addr = static_cast<unsigned>(__cvta_generic_to_shared(dst));
unsigned long long gmem_addr = __cvta_generic_to_global(src);
asm volatile("cp.async.cg.shared.global [%0], [%1], 16;\n" :: "r"(smem_addr), "l"(gmem_addr));
}
asm volatile("cp.async.commit_group;\n" ::);
asm volatile("cp.async.wait_all;\n" ::);
const int tail_start = (K_bytes & ~15);
for (int idx = tail_start + threadIdx.x; idx < K_bytes; idx += blockDim.x) {
sh_B[idx] = B[b_row_base + (int64_t)idx * b_stride_k];
}
} else
#endif
{
for (int idx = threadIdx.x; idx < K_bytes; idx += blockDim.x) {
const int64_t g_idx = b_row_base + (int64_t)idx * b_stride_k;
sh_B[idx] = B[g_idx];
}
}
for (int idx = threadIdx.x; idx < K_scales; idx += blockDim.x) {
const int64_t g_idx = sfb_row_base + (int64_t)idx * sfb_stride_k;
sh_SFB[idx] = decode_fp8_e4m3(SFB[g_idx]);
}
__syncthreads();
if (m >= M) {
return;
}
const int64_t a_row_base = (int64_t)m * a_stride_m + (int64_t)l * a_stride_l;
const int64_t sfa_row_base = (int64_t)m * sfa_stride_m + (int64_t)l * sfa_stride_l;
float acc = 0.0f;
for (int g = 0; g < K_scales; ++g) {
const int64_t sfa_idx = sfa_row_base + (int64_t)g * sfa_stride_k;
const float scale_a = decode_fp8_e4m3(SFA[sfa_idx]);
const float scale_b = sh_SFB[g];
const float block_scale = scale_a * scale_b;
if (block_scale == 0.0f) {
continue;
}
float group_sum = 0.0f;
const int byte_start = g * 8;
#pragma unroll 8
for (int i = 0; i < 8; ++i) {
const int byte_idx = byte_start + i;
const int64_t a_idx = a_row_base + (int64_t)byte_idx * a_stride_k;
const uint8_t a_byte = A[a_idx];
const uint8_t b_byte = sh_B[byte_idx];
const float a0 = decode_fp4_e2m1(a_byte & 0xF);
const float a1 = decode_fp4_e2m1((a_byte >> 4) & 0xF);
const float b0 = decode_fp4_e2m1(b_byte & 0xF);
const float b1 = decode_fp4_e2m1((b_byte >> 4) & 0xF);
group_sum += a0 * b0;
group_sum += a1 * b1;
}
acc += group_sum * block_scale;
}
const int64_t c_idx = (int64_t)m * L + l;
C[c_idx] = __float2half_rn(acc);
}
#if ENABLE_WMMA_KERNEL
__global__ void bmv_kernel_tensor(
const uint8_t* __restrict__ A,
const uint8_t* __restrict__ B,
const uint8_t* __restrict__ SFA,
const uint8_t* __restrict__ SFB,
__half* __restrict__ C,
int M, int M_B, int K_bytes, int L,
int64_t a_stride_m, int64_t a_stride_k, int64_t a_stride_l,
int64_t b_stride_m, int64_t b_stride_k, int64_t b_stride_l,
int64_t sfa_stride_m, int64_t sfa_stride_k, int64_t sfa_stride_l,
int64_t sfb_stride_m, int64_t sfb_stride_k, int64_t sfb_stride_l)
{
#if defined(__CUDA_ARCH__) && __CUDA_ARCH__ >= 800
const int tile_m = blockIdx.x * CTA_M;
const int tile_l = blockIdx.y * CTA_N;
if (tile_m >= M || tile_l >= L) {
return;
}
const int warp_id = threadIdx.x / warpSize;
const int lane_id = threadIdx.x % warpSize;
const int K_vals = K_bytes * 2;
const int K_scales = K_bytes / 8;
const int m_b = 0;
extern __shared__ uint8_t shared_raw[];
half* sh_A = reinterpret_cast<half*>(shared_raw);
half* sh_B = sh_A + CTA_M * WMMA_K;
float* sh_out = reinterpret_cast<float*>(sh_B + CTA_N * WMMA_K);
wmma::fragment<wmma::accumulator, WMMA_M, WMMA_N, WMMA_K, float> acc_frag;
wmma::fill_fragment(acc_frag, 0.0f);
for (int k_base = 0; k_base < K_vals; k_base += WMMA_K) {
// Stage A tile
for (int idx = threadIdx.x; idx < CTA_M * WMMA_K; idx += blockDim.x) {
const int row = idx / WMMA_K;
const int k_inner = idx % WMMA_K;
const int global_m = tile_m + row;
const int global_k = k_base + k_inner;
float val = 0.0f;
if (global_m < M && global_k < K_vals) {
const int64_t a_row_base = (int64_t)global_m * a_stride_m + (int64_t)tile_l * a_stride_l;
const int64_t sfa_row_base = (int64_t)global_m * sfa_stride_m + (int64_t)tile_l * sfa_stride_l;
const int g = global_k / 16;
if (g < K_scales) {
const float scale_a = decode_fp8_e4m3(SFA[sfa_row_base + (int64_t)g * sfa_stride_k]);
const int byte_index = global_k >> 1;
const uint8_t byte_val = A[a_row_base + (int64_t)byte_index * a_stride_k];
const bool high = (global_k & 1);
const uint8_t nibble = high ? ((byte_val >> 4) & 0xF) : (byte_val & 0xF);
val = decode_fp4_e2m1(nibble) * scale_a;
}
}
sh_A[idx] = __float2half(val);
}
// Stage B tile
const int64_t b_row_base = (int64_t)m_b * b_stride_m + (int64_t)tile_l * b_stride_l;
const int64_t sfb_row_base = (int64_t)m_b * sfb_stride_m + (int64_t)tile_l * sfb_stride_l;
for (int idx = threadIdx.x; idx < CTA_N * WMMA_K; idx += blockDim.x) {
const int col = idx / WMMA_K;
const int k_inner = idx % WMMA_K;
const int global_l = tile_l + col;
const int global_k = k_base + k_inner;
float val = 0.0f;
if (global_l < L && global_k < K_vals) {
const int g = global_k / 16;
if (g < K_scales) {
const float scale_b = decode_fp8_e4m3(SFB[sfb_row_base + (int64_t)g * sfb_stride_k]);
const int byte_index = global_k >> 1;
const uint8_t byte_val = B[b_row_base + (int64_t)byte_index * b_stride_k];
const bool high = (global_k & 1);
const uint8_t nibble = high ? ((byte_val >> 4) & 0xF) : (byte_val & 0xF);
val = decode_fp4_e2m1(nibble) * scale_b;
}
}
sh_B[idx] = __float2half(val);
}
__syncthreads();
if (warp_id < (CTA_M / WMMA_M)) {
const half* warp_A = sh_A + warp_id * WMMA_M * WMMA_K;
wmma::fragment<wmma::matrix_a, WMMA_M, WMMA_N, WMMA_K, half, wmma::row_major> a_frag;
wmma::fragment<wmma::matrix_b, WMMA_M, WMMA_N, WMMA_K, half, wmma::col_major> b_frag;
wmma::load_matrix_sync(a_frag, warp_A, WMMA_K);
wmma::load_matrix_sync(b_frag, sh_B, WMMA_K);
wmma::mma_sync(acc_frag, a_frag, b_frag, acc_frag);
}
__syncthreads();
}
if (warp_id < (CTA_M / WMMA_M)) {
float* warp_out = sh_out + warp_id * WMMA_M * CTA_N;
wmma::store_matrix_sync(warp_out, acc_frag, CTA_N, wmma::mem_row_major);
}
__syncthreads();
for (int idx = threadIdx.x; idx < CTA_M * CTA_N; idx += blockDim.x) {
const int row = idx / CTA_N;
const int col = idx % CTA_N;
const int global_m = tile_m + row;
const int global_l = tile_l + col;
if (global_m < M && global_l < L) {
const float val = sh_out[idx];
C[global_m * L + global_l] = __float2half(val);
}
}
#else
(void)A; (void)B; (void)SFA; (void)SFB; (void)C;
(void)M; (void)M_B; (void)K_bytes; (void)L;
(void)a_stride_m; (void)a_stride_k; (void)a_stride_l;
(void)b_stride_m; (void)b_stride_k; (void)b_stride_l;
(void)sfa_stride_m; (void)sfa_stride_k; (void)sfa_stride_l;
(void)sfb_stride_m; (void)sfb_stride_k; (void)sfb_stride_l;
#endif
}
#endif // ENABLE_WMMA_KERNEL
#if NVFP4_ENABLE_TCGEN
__global__ void bmv_kernel_tcgen(
const uint8_t* __restrict__ A,
const uint8_t* __restrict__ B,
const uint8_t* __restrict__ SFA,
const uint8_t* __restrict__ SFB,
__half* __restrict__ C,
int M, int M_B, int K_bytes, int L,
int64_t a_stride_m, int64_t a_stride_k, int64_t a_stride_l,
int64_t b_stride_m, int64_t b_stride_k, int64_t b_stride_l,
int64_t sfa_stride_m, int64_t sfa_stride_k, int64_t sfa_stride_l,
int64_t sfb_stride_m, int64_t sfb_stride_k, int64_t sfb_stride_l)
{
#if defined(__CUDA_ARCH__) && defined(__CUDA_ARCH_FAMILY_SPECIFIC__)
constexpr uint32_t TMEM_COLS_ACC = 32;
constexpr uint32_t TMEM_COLS_A = 32;
constexpr uint32_t TMEM_COLS_B = 32;
constexpr uint32_t TMEM_COLS_SFA = 32;
constexpr uint32_t TMEM_COLS_SFB = 32;
__shared__ uint32_t sh_tmem_acc;
__shared__ uint32_t sh_tmem_a;
__shared__ uint32_t sh_tmem_b;
__shared__ uint32_t sh_tmem_sfa;
__shared__ uint32_t sh_tmem_sfb;
__shared__ uint64_t sh_mbarrier;
const int tile_m = blockIdx.x * TCGEN_CTA_M;
const int tile_l = blockIdx.y;
if (tile_m >= M || tile_l >= L) {
return;
}
const int warp_id = threadIdx.x / warpSize;
const int lane_id = threadIdx.x % warpSize;
if (warp_id >= TCGEN_CTA_WARPS) {
return;
}
const int lane_group = lane_id >> 2;
const int lane_tid4 = lane_id & 3;
const int warp_row_base = tile_m + warp_id * WMMA_M;
if (warp_row_base >= M) {
return;
}
const int K_vals = K_bytes * 2;
const int K_scales = K_bytes / 8;
const int m_b = 0;
const int b_row = 0;
if (threadIdx.x == 0) {
tcgen_alloc_cols(&sh_tmem_acc, TMEM_COLS_ACC);
tcgen_alloc_cols(&sh_tmem_a, TMEM_COLS_A);
tcgen_alloc_cols(&sh_tmem_b, TMEM_COLS_B);
tcgen_alloc_cols(&sh_tmem_sfa, TMEM_COLS_SFA);
tcgen_alloc_cols(&sh_tmem_sfb, TMEM_COLS_SFB);
}
__syncthreads();
const uint32_t tmem_acc_addr = sh_tmem_acc;
const uint32_t tmem_a_addr = sh_tmem_a;
const uint32_t tmem_b_addr = sh_tmem_b;
const uint32_t tmem_sfa_addr = sh_tmem_sfa;
const uint32_t tmem_sfb_addr = sh_tmem_sfb;
(void)tmem_acc_addr;
if (threadIdx.x == 0) {
mbarrier_init_cta(&sh_mbarrier, 0);
}
__syncthreads();
uint32_t mbarrier_parity = 0;
extern __shared__ uint8_t shared_tiles[];
struct StageBuffers {
uint8_t* A;
uint8_t* B;
uint8_t* SFA;
uint8_t* SFB;
};
const size_t bytes_A_row = WMMA_K / 2;
const size_t bytes_A_tile = (size_t)TCGEN_CTA_M * bytes_A_row;
const size_t bytes_B_tile = WMMA_K / 2;
const size_t bytes_SFA_tile = (size_t)TCGEN_CTA_M * 4;
const size_t bytes_SFB_tile = 4;
const size_t stage_bytes = bytes_A_tile + bytes_B_tile + bytes_SFA_tile + bytes_SFB_tile;
auto stage_buffers = [&](int idx) {
StageBuffers buf;
uint8_t* base = shared_tiles + idx * stage_bytes;
buf.A = base;
buf.B = buf.A + bytes_A_tile;
buf.SFA = buf.B + bytes_B_tile;
buf.SFB = buf.SFA + bytes_SFA_tile;
return buf;
};
auto make_desc = [&](const void* ptr, uint32_t lead_bytes, uint32_t stride_bytes) {
return make_smem_matrix_desc(smem_ptr(ptr), lead_bytes, stride_bytes, true);
};
const int64_t b_row_base = (int64_t)m_b * b_stride_m + (int64_t)tile_l * b_stride_l;
const int64_t sfb_row_base = (int64_t)m_b * sfb_stride_m + (int64_t)tile_l * sfb_stride_l;
auto load_stage = [&](const StageBuffers& buf, int k_base_panel) {
const int byte_offset = k_base_panel >> 1;
for (int idx = threadIdx.x; idx < bytes_A_tile; idx += blockDim.x) {
const int row_local = idx / bytes_A_row;
const int byte_in_row = idx % bytes_A_row;
const int global_row = tile_m + row_local;
const int global_byte = byte_offset + byte_in_row;
uint8_t val = 0;
if (global_row < M && global_byte < K_bytes) {
const int64_t a_offset = (int64_t)global_row * a_stride_m +
(int64_t)global_byte * a_stride_k +
(int64_t)tile_l * a_stride_l;
val = A[a_offset];
}
buf.A[idx] = val;
}
#if CP_ASYNC_SUPPORTED
if (b_stride_k == 1) {
const int chunk = 16;
const int copy_bytes = bytes_B_tile & ~ (chunk - 1);
for (int off = threadIdx.x * chunk; off < copy_bytes; off += blockDim.x * chunk) {
const int global_byte = byte_offset + off;
const uint32_t dst = smem_ptr(buf.B + off);
const unsigned long long src = __cvta_generic_to_global(B + b_row_base + global_byte);
asm volatile("cp.async.cg.shared.global [%0], [%1], 16;\n" :: "r"(dst), "l"(src));
}
asm volatile("cp.async.commit_group;\n" ::);
asm volatile("cp.async.wait_all;\n" ::);
for (int off = copy_bytes + threadIdx.x; off < bytes_B_tile; off += blockDim.x) {
const int global_byte = byte_offset + off;
uint8_t val = 0;
if (global_byte < K_bytes) {
const int64_t b_offset = b_row_base + (int64_t)global_byte * b_stride_k;
val = B[b_offset];
}
buf.B[off] = val;
}
} else
#endif
{
for (int idx = threadIdx.x; idx < bytes_B_tile; idx += blockDim.x) {
const int global_byte = byte_offset + idx;
uint8_t val = 0;
if (global_byte < K_bytes) {
const int64_t b_offset = b_row_base + (int64_t)global_byte * b_stride_k;
val = B[b_offset];
}
buf.B[idx] = val;
}
}
const int scale_offset = k_base_panel >> 4;
for (int idx = threadIdx.x; idx < bytes_SFA_tile; idx += blockDim.x) {
const int row_local = idx / 4;
const int elem = idx & 3;
const int global_row = tile_m + row_local;
const int scale_idx = scale_offset + elem;
uint8_t val = 0;
if (global_row < M && scale_idx < K_scales) {
const int64_t sfa_offset = (int64_t)global_row * sfa_stride_m +
(int64_t)scale_idx * sfa_stride_k +
(int64_t)tile_l * sfa_stride_l;
val = SFA[sfa_offset];
}
buf.SFA[idx] = val;
}
for (int idx = threadIdx.x; idx < bytes_SFB_tile; idx += blockDim.x) {
const int scale_idx = scale_offset + idx;
uint8_t val = 0;
if (scale_idx < K_scales) {
const int64_t sfb_offset = sfb_row_base + (int64_t)scale_idx * sfb_stride_k;
val = SFB[sfb_offset];
}
buf.SFB[idx] = val;
}
};
StageBuffers stage0 = stage_buffers(0);
StageBuffers stage1 = stage_buffers(1);
load_stage(stage0, 0);
__syncthreads();
int stage_idx = 0;
float d0 = 0.0f, d1 = 0.0f, d2 = 0.0f, d3 = 0.0f;
const uint16_t bidA = 0, tidA = 0;
const uint16_t bidB = 0, tidB = 0;
#if NVFP4_EXPERIMENTAL_TCGEN_MMA
const uint32_t instr_desc_tcgen = make_mxf4nvf4_instr_desc(TCGEN_CTA_M, WMMA_N, true, 0, 0);
#endif
for (int k_base = 0; k_base < K_vals; k_base += WMMA_K) {
const StageBuffers& stage = (stage_idx == 0) ? stage0 : stage1;
uint32_t a0 = 0u, a1 = 0u, a2 = 0u, a3 = 0u;
uint32_t b0 = 0u, b1 = 0u;
uint32_t scaleAData = 0u;
uint32_t scaleBData = 0u;
if (threadIdx.x == 0) {
const uint64_t descA = make_desc(stage.A, bytes_A_row, bytes_A_tile);
const uint64_t descB = make_desc(stage.B, bytes_B_tile, bytes_B_tile);
const uint64_t descSFA = make_desc(stage.SFA, 4u, bytes_SFA_tile);
const uint64_t descSFB = make_desc(stage.SFB, 4u, bytes_SFB_tile);
tcgen_cp_async(tmem_a_addr, descA, 0);
tcgen_cp_async(tmem_b_addr, descB, 0);
tcgen_cp_async(tmem_sfa_addr, descSFA, 0);
tcgen_cp_async(tmem_sfb_addr, descSFB, 0);
}
__syncthreads();
if (threadIdx.x == 0) {
tcgen_commit_mbarrier(&sh_mbarrier);
}
mbarrier_wait_parity_cta(&sh_mbarrier, mbarrier_parity);
mbarrier_parity ^= 1;
#if NVFP4_EXPERIMENTAL_TCGEN_MMA
if (threadIdx.x == 0) {
const uint64_t descB_tc = make_desc(stage.B, bytes_B_tile, bytes_B_tile);
const uint32_t enable_flag = (k_base == 0) ? 0u : 1u;
asm volatile(
"{\n\t"
".reg .pred p_enable;\n\t"
"setp.ne.u32 p_enable, %6, 0;\n\t"
"tcgen05.mma.cta_group::1.kind::mxf4nvf4.block_scale.scale_vec::4X "
"[%0], [%1], %2, %3, [%4], [%5], p_enable;\n\t"
"}\n"
:
: "r"(tmem_acc_addr),
"r"(tmem_a_addr),
"l"(descB_tc),
"r"(instr_desc_tcgen),
"r"(tmem_sfa_addr),
"r"(tmem_sfb_addr),
"r"(enable_flag)
: "memory"
);
}
#endif
#pragma unroll
for (int elem = 0; elem < 32; ++elem) {
const int row_local = frag_a_row(lane_group, elem);
const int col_local = frag_a_col(lane_tid4, elem);
const int row_cta = warp_id * WMMA_M + row_local;
const uint8_t nibble = load_stage_a_nibble(stage.A, row_cta, col_local);
uint32_t* dst = (elem < 8) ? &a0 : (elem < 16) ? &a1 : (elem < 24) ? &a2 : &a3;
const int shift = (elem & 7) * 4;
*dst |= uint32_t(nibble & 0xF) << shift;
}
#pragma unroll
for (int elem = 0; elem < 16; ++elem) {
const int row_local = frag_b_row(lane_tid4, elem);
const uint8_t nibble = load_stage_b_nibble(stage.B, row_local);
uint32_t* dst = (elem < 8) ? &b0 : &b1;
const int shift = (elem & 7) * 4;
*dst |= uint32_t(nibble & 0xF) << shift;
}
if (lane_tid4 <= 1) {
const int row_local = lane_group + (lane_tid4 ? 8 : 0);
const int row_cta = warp_id * WMMA_M + row_local;
uint8_t s0 = load_stage_sfa_byte(stage.SFA, row_cta, 0);
uint8_t s1 = load_stage_sfa_byte(stage.SFA, row_cta, 1);
uint8_t s2 = load_stage_sfa_byte(stage.SFA, row_cta, 2);
uint8_t s3 = load_stage_sfa_byte(stage.SFA, row_cta, 3);
scaleAData = uint32_t(s0) | (uint32_t(s1) << 8) | (uint32_t(s2) << 16) | (uint32_t(s3) << 24);
}
if (lane_group == 0 && lane_tid4 == 0) {
uint8_t s0 = load_stage_sfb_byte(stage.SFB, 0);
uint8_t s1 = load_stage_sfb_byte(stage.SFB, 1);
uint8_t s2 = load_stage_sfb_byte(stage.SFB, 2);
uint8_t s3 = load_stage_sfb_byte(stage.SFB, 3);
scaleBData = uint32_t(s0) | (uint32_t(s1) << 8) | (uint32_t(s2) << 16) | (uint32_t(s3) << 24);
}
float c0 = d0, c1 = d1, c2 = d2, c3 = d3;
asm volatile(
"mma.sync.aligned.m16n8k64.row.col.kind::mxf4nvf4.block_scale.scale_vec::4X "
".f32.e2m1.e2m1.f32.ue4m3 "
"{%0,%1,%2,%3}, "
"{%4,%5,%6,%7}, "
"{%8,%9}, "
"{%10,%11,%12,%13}, "
"%14, {%15,%16}, "
"%17, {%18,%19};\n"
: "+f"(d0), "+f"(d1), "+f"(d2), "+f"(d3)
: "r"(a0), "r"(a1), "r"(a2), "r"(a3),
"r"(b0), "r"(b1),
"f"(c0), "f"(c1), "f"(c2), "f"(c3),
"r"(scaleAData), "h"(bidA), "h"(tidA),
"r"(scaleBData), "h"(bidB), "h"(tidB)
);
stage_idx ^= 1;
const int next_k = k_base + WMMA_K;
if (next_k < K_vals) {
const StageBuffers& next_stage = (stage_idx == 0) ? stage0 : stage1;
load_stage(next_stage, next_k);
}
__syncthreads();
}
#if NVFP4_EXPERIMENTAL_TCGEN_MMA
const float wmma_ref0 = d0;
const float wmma_ref1 = d1;
const float wmma_ref2 = d2;
const float wmma_ref3 = d3;
__syncthreads();
if (threadIdx.x == 0) {
tcgen_commit_mbarrier(&sh_mbarrier);
}
mbarrier_wait_parity_cta(&sh_mbarrier, mbarrier_parity);
mbarrier_parity ^= 1;
const uint32_t warp_lane_base = tmem_add_lane_offset(tmem_acc_addr, static_cast<uint32_t>(warp_id * 32));
float tm_d0, tm_d1, tm_d2, tm_d3;
asm volatile(
"tcgen05.ld.sync.aligned.16x64b.x4.b32 {%0,%1,%2,%3}, [%4];\n"
: "=f"(tm_d0), "=f"(tm_d1), "=f"(tm_d2), "=f"(tm_d3)
: "r"(warp_lane_base)
: "memory"
);
d0 = tm_d0;
d1 = tm_d1;
d2 = tm_d2;
d3 = tm_d3;
if (blockIdx.x == 0 && blockIdx.y == 0) {
const int base_idx = threadIdx.x * 4;
g_debug_tm_frag[base_idx + 0] = tm_d0;
g_debug_tm_frag[base_idx + 1] = tm_d1;
g_debug_tm_frag[base_idx + 2] = tm_d2;
g_debug_tm_frag[base_idx + 3] = tm_d3;
g_debug_ref_frag[base_idx + 0] = wmma_ref0;
g_debug_ref_frag[base_idx + 1] = wmma_ref1;
g_debug_ref_frag[base_idx + 2] = wmma_ref2;
g_debug_ref_frag[base_idx + 3] = wmma_ref3;
}
#endif
float row_sum_top = 0.0f;
float row_sum_bottom = 0.0f;
#pragma unroll
for (int acc_idx = 0; acc_idx < 4; ++acc_idx) {
const float value = (acc_idx == 0) ? d0 : (acc_idx == 1) ? d1 : (acc_idx == 2) ? d2 : d3;
const int row_local = frag_acc_row(lane_group, acc_idx);
if (row_local == lane_group) {
row_sum_top += value;
} else {
row_sum_bottom += value;
}
}
for (int offset = 2; offset > 0; offset >>= 1) {
row_sum_top += __shfl_xor_sync(0xFFFFFFFF, row_sum_top, offset, 4);
row_sum_bottom += __shfl_xor_sync(0xFFFFFFFF, row_sum_bottom, offset, 4);
}
if (lane_tid4 == 0) {
const int global_col = tile_l;
const int row0 = warp_row_base + lane_group;
if (row0 < M && global_col < L) {
const int64_t out_idx = (int64_t)row0 * L + global_col;
C[out_idx] = __float2half_rn(row_sum_top);
}
const int row1 = warp_row_base + lane_group + 8;
if (row1 < M && global_col < L) {
const int64_t out_idx = (int64_t)row1 * L + global_col;
C[out_idx] = __float2half_rn(row_sum_bottom);
}
}
__syncthreads();
if (threadIdx.x == 0) {
tcgen_dealloc_cols(tmem_sfb_addr, TMEM_COLS_SFB);
tcgen_dealloc_cols(tmem_sfa_addr, TMEM_COLS_SFA);
tcgen_dealloc_cols(tmem_b_addr, TMEM_COLS_B);
tcgen_dealloc_cols(tmem_a_addr, TMEM_COLS_A);
tcgen_dealloc_cols(tmem_acc_addr, TMEM_COLS_ACC);
}
#else
(void)A; (void)B; (void)SFA; (void)SFB; (void)C;
(void)M; (void)M_B; (void)K_bytes; (void)L;
(void)a_stride_m; (void)a_stride_k; (void)a_stride_l;
(void)b_stride_m; (void)b_stride_k; (void)b_stride_l;
(void)sfa_stride_m; (void)sfa_stride_k; (void)sfa_stride_l;
(void)sfb_stride_m; (void)sfb_stride_k; (void)sfb_stride_l;
#endif
}
#endif // NVFP4_ENABLE_TCGEN
void bmv_launcher(
torch::Tensor A,
torch::Tensor B,
torch::Tensor SFA,
torch::Tensor SFB,
torch::Tensor C)
{
TORCH_CHECK(A.dim() == 3, "A must be 3D");
TORCH_CHECK(B.dim() == 3, "B must be 3D");
TORCH_CHECK(SFA.dim() == 3, "SFA must be 3D");
TORCH_CHECK(SFB.dim() == 3, "SFB must be 3D");
const int64_t M = A.size(0);
const int64_t K_bytes = A.size(1);
const int64_t L = A.size(2);
const int64_t M_B = B.size(0);
TORCH_CHECK(M_B >= 1, "B must have at least one row");
TORCH_CHECK(B.size(1) == K_bytes && B.size(2) == L, "B shape mismatch");
TORCH_CHECK(SFA.size(0) == M && SFA.size(1) == K_bytes / 8 && SFA.size(2) == L, "SFA shape mismatch");
TORCH_CHECK(SFB.size(0) >= 1 && SFB.size(1) == K_bytes / 8 && SFB.size(2) == L, "SFB shape mismatch");
const auto a_stride = A.strides();
const auto b_stride = B.strides();
const auto sfa_stride = SFA.strides();
const auto sfb_stride = SFB.strides();
cudaStream_t stream = c10::cuda::getCurrentCUDAStream().stream();
auto align_up = [](size_t value, size_t alignment) {
return (value + alignment - 1) & ~(alignment - 1);
};
const cudaDeviceProp* device_prop = at::cuda::getCurrentDeviceProperties();
const KernelPath override = kernel_path_override_from_env();
const SmallNMode small_n_mode = small_n_mode_override_from_env();
const KernelPath selected = choose_kernel_path(device_prop, override, small_n_mode, static_cast<int>(K_bytes), static_cast<int>(L));
#if NVFP4_EXPERIMENTAL_TCGEN_MMA
const bool capture_tmem_debug = capture_tmem_debug_enabled();
if (capture_tmem_debug) {
printf("[nvfp4] selected kernel path = %d\n", static_cast<int>(selected));
}
#else
const bool capture_tmem_debug = false;
#endif
#if !NVFP4_ENABLE_TCGEN
TORCH_CHECK(
selected != KernelPath::TCGEN,
"TCGEN kernel requested but this build was compiled without NVFP4 tensor-core support. "
"Set NVFP4_ENABLE_TCGEN_BUILD=1 before importing submission.py to enable it.");
#endif
switch (selected) {
#if NVFP4_ENABLE_TCGEN
case KernelPath::TCGEN: {
TORCH_CHECK(
arch_supports_tcgen(device_prop),
"TCGEN path selected but device does not support it.");
const int threads = TCGEN_THREADS;
const dim3 grid_tcgen(
static_cast<unsigned int>((M + TCGEN_CTA_M - 1) / TCGEN_CTA_M),
static_cast<unsigned int>(L)
);
const size_t bytes_A_tile = static_cast<size_t>(TCGEN_CTA_M) * (WMMA_K / 2);
const size_t bytes_B_tile = WMMA_K / 2;
const size_t bytes_SFA_tile = static_cast<size_t>(TCGEN_CTA_M) * 4;
const size_t bytes_SFB_tile = 4;
const size_t stage_bytes = bytes_A_tile + bytes_B_tile + bytes_SFA_tile + bytes_SFB_tile;
const size_t shared_tcgen = stage_bytes * 2;
#if NVFP4_EXPERIMENTAL_TCGEN_MMA
if (capture_tmem_debug) {
const size_t debug_bytes = static_cast<size_t>(TCGEN_THREADS) * 4 * sizeof(float);
cudaMemsetAsync(g_debug_tm_frag, 0, debug_bytes, stream);
cudaMemsetAsync(g_debug_ref_frag, 0, debug_bytes, stream);
}
#endif
bmv_kernel_tcgen<<<grid_tcgen, threads, shared_tcgen, stream>>>(
A.data_ptr<uint8_t>(),
B.data_ptr<uint8_t>(),
SFA.data_ptr<uint8_t>(),
SFB.data_ptr<uint8_t>(),
reinterpret_cast<__half*>(C.data_ptr<at::Half>()),
static_cast<int>(M),
static_cast<int>(M_B),
static_cast<int>(K_bytes),
static_cast<int>(L),
a_stride[0], a_stride[1], a_stride[2],
b_stride[0], b_stride[1], b_stride[2],
sfa_stride[0], sfa_stride[1], sfa_stride[2],
sfb_stride[0], sfb_stride[1], sfb_stride[2]
);
break;
}
#endif
#if ENABLE_WMMA_KERNEL
case KernelPath::WMMA: {
TORCH_CHECK(
arch_supports_wmma(device_prop),
"WMMA tensor path selected but device does not support it.");
const dim3 grid_tensor(
static_cast<unsigned int>((M + CTA_M - 1) / CTA_M),
static_cast<unsigned int>((L + CTA_N - 1) / CTA_N)
);
const size_t shared_tensor =
(static_cast<size_t>(CTA_M) * WMMA_K +
static_cast<size_t>(CTA_N) * WMMA_K) * sizeof(half) +
static_cast<size_t>(CTA_M) * CTA_N * sizeof(float);
bmv_kernel_tensor<<<grid_tensor, 256, shared_tensor, stream>>>(
A.data_ptr<uint8_t>(),
B.data_ptr<uint8_t>(),
SFA.data_ptr<uint8_t>(),
SFB.data_ptr<uint8_t>(),
reinterpret_cast<__half*>(C.data_ptr<at::Half>()),
static_cast<int>(M),
static_cast<int>(M_B),
static_cast<int>(K_bytes),
static_cast<int>(L),
a_stride[0], a_stride[1], a_stride[2],
b_stride[0], b_stride[1], b_stride[2],
sfa_stride[0], sfa_stride[1], sfa_stride[2],
sfb_stride[0], sfb_stride[1], sfb_stride[2]
);
break;
}
#endif
case KernelPath::Scalar:
default: {
const int threads = 256;
const dim3 grid_scalar(
static_cast<unsigned int>((M + threads - 1) / threads),
static_cast<unsigned int>(L)
);
const size_t shared_scalar = align_up(static_cast<size_t>(K_bytes), size_t(16)) +
static_cast<size_t>(K_bytes / 8) * sizeof(float);
bmv_kernel_scalar<<<grid_scalar, threads, shared_scalar, stream>>>(
A.data_ptr<uint8_t>(),
B.data_ptr<uint8_t>(),
SFA.data_ptr<uint8_t>(),
SFB.data_ptr<uint8_t>(),
reinterpret_cast<__half*>(C.data_ptr<at::Half>()),
static_cast<int>(M),
static_cast<int>(M_B),
static_cast<int>(K_bytes),
static_cast<int>(L),
a_stride[0], a_stride[1], a_stride[2],
b_stride[0], b_stride[1], b_stride[2],
sfa_stride[0], sfa_stride[1], sfa_stride[2],
sfb_stride[0], sfb_stride[1], sfb_stride[2]
);
break;
}
}
cudaError_t err = cudaGetLastError();
if (err != cudaSuccess) {
TORCH_CHECK(false, "CUDA kernel launch failed: ", cudaGetErrorString(err));
}
#if NVFP4_EXPERIMENTAL_TCGEN_MMA
if (selected == KernelPath::TCGEN && capture_tmem_debug) {
const size_t debug_elems = static_cast<size_t>(TCGEN_THREADS) * 4;
const size_t debug_bytes = debug_elems * sizeof(float);
std::vector<float> debug_tm(debug_elems, 0.0f);
std::vector<float> debug_ref(debug_elems, 0.0f);
cudaMemcpyFromSymbolAsync(debug_tm.data(), g_debug_tm_frag, debug_bytes, 0, cudaMemcpyDeviceToHost, stream);
cudaMemcpyFromSymbolAsync(debug_ref.data(), g_debug_ref_frag, debug_bytes, 0, cudaMemcpyDeviceToHost, stream);
cudaStreamSynchronize(stream);
printf("==== TMEM debug (CTA 0, warp 0) ====\n");
for (int thread = 0; thread < TCGEN_THREADS; ++thread) {
const int warp = thread / 32;
if (warp > 0) {
continue;
}
const int lane = thread % 32;
const int group = lane >> 2;
const int tid4 = lane & 3;
for (int acc_idx = 0; acc_idx < 4; ++acc_idx) {
const float tm_val = debug_tm[thread * 4 + acc_idx];
const float ref_val = debug_ref[thread * 4 + acc_idx];
if (tm_val == 0.0f && ref_val == 0.0f) {
continue;
}
const int row_local = (acc_idx < 2) ? group : (group + 8);
const int col_local = tid4 * 2 + (acc_idx & 1);
const float diff = std::fabs(tm_val - ref_val);
printf("lane=%d acc=%d row=%d col=%d tm=%f ref=%f diff=%f\n",
lane, acc_idx, row_local, col_local, tm_val, ref_val, diff);
}
}
printf("==== TMEM debug end ====\n");
}
#endif
}
"""
# Build the extension
extra_cuda_cflags = [
"-O3",
"-std=c++17",
"-gencode=arch=compute_80,code=sm_80",
"-gencode=arch=compute_100,code=sm_100",
"-U__CUDA_NO_HALF_OPERATORS__",
"-U__CUDA_NO_HALF_CONVERSIONS__",
"-U__CUDA_NO_BFLOAT16_CONVERSIONS__",
"--expt-relaxed-constexpr",
"-use_fast_math",
f"-DNVFP4_ENABLE_TCGEN={1 if enable_tcgen_build else 0}",
f"-DNVFP4_EXPERIMENTAL_TCGEN_MMA={1 if experimental_tcgen else 0}",
]
if os.environ.get("NVFP4_CAPTURE_TMEM_DEBUG"):
extra_cuda_cflags.extend(["-keep", "--keep-dir=/tmp/nvcc_keep"])
if enable_tcgen_build:
generic_flag = "-gencode=arch=compute_100,code=sm_100"
extra_cuda_cflags = [flag for flag in extra_cuda_cflags if flag != generic_flag]
extra_cuda_cflags.insert(4, "-gencode=arch=compute_100a,code=sm_100a")
_EXT = load_inline(
name="bmv_nvfp4_b200_ext",
cpp_sources=[CPP_SRC],
cuda_sources=[CUDA_SRC],
functions=["bmv_launcher"],
with_cuda=True,
extra_cflags=["-O3", "-std=c++17"],
extra_cuda_cflags=extra_cuda_cflags,
verbose=False,
)
return _EXT
def _convert_to_uint8(tensor):
"""Convert FP4/FP8 dtype tensors to uint8 tensor containing raw bytes."""
if hasattr(torch, 'float4_e2m1fn') and tensor.dtype == torch.float4_e2m1fn:
return tensor.view(torch.uint8)
elif hasattr(torch, 'float4_e2m1fn_x2') and str(tensor.dtype) == 'torch.float4_e2m1fn_x2':
return tensor.view(torch.uint8)
elif tensor.dtype in [torch.float8_e4m3fn, torch.float8_e5m2]:
return tensor.view(torch.uint8)
elif tensor.dtype == torch.uint8:
return tensor
else:
try:
return tensor.view(torch.uint8)
except:
raise ValueError(f"Cannot convert dtype {tensor.dtype} to uint8")
@torch.inference_mode()
def custom_kernel(tensors):
"""
NVFP4 batched matrix-vector multiplication with FP8 block scales.
Always returns a NEW tensor of shape [M, 1, L] in FP16.
Expected input: (a, b, sfa, sfb, c)
a: [M, K_bytes, L], FP4 packed (2 FP4 vals/byte)
b: [M_B, K_bytes, L], FP4 packed where M_B can be any value
sfa:[M, K_bytes/8, L], FP8 scales
sfb:[M_B, K_bytes/8, L], FP8 scales
c: Output tensor (any shape - only used for device)
Returns: NEW tensor of shape [M, 1, L] in FP16
"""
if len(tensors) >= 5:
a, b, sfa, sfb, c = tensors[:5]
else:
raise ValueError(f"Expected at least 5 tensors, got {len(tensors)}")
# Get dimensions
M, K_bytes, L = a.shape
M_B = b.shape[0]
# Validate compatible K and L dimensions
assert b.shape[1] == K_bytes and b.shape[2] == L, f"B shape {b.shape} incompatible with A shape {a.shape}"
# Validate scale shapes
K_scales = K_bytes // 8
assert sfa.shape == (M, K_scales, L), f"SFA shape {sfa.shape} should be ({M}, {K_scales}, {L})"
assert sfb.shape == (M_B, K_scales, L), f"SFB shape {sfb.shape} should be ({M_B}, {K_scales}, {L})"
# Validate K alignment
assert (K_bytes * 2) % 16 == 0, f"K must be multiple of 16 FP4 elements"
# Convert to uint8 raw bytes without altering layout (preserve strides)
a = _convert_to_uint8(a)
b = _convert_to_uint8(b)
sfa = _convert_to_uint8(sfa)
sfb = _convert_to_uint8(sfb)
# Create NEW output tensor with correct shape [M, 1, L] in FP16
output = torch.zeros(M, 1, L, dtype=torch.float16, device=c.device).contiguous()
# Call CUDA kernel
ext = _get_ext()
ext.bmv_launcher(a, b, sfa, sfb, output)
# Ensure completion
torch.cuda.synchronize()
return output
scrolls · 1573 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