submission 104335
John Hahn · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 759 lines, June 9 Researcher Reciprocity License v1.0.
mid8.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-nvfp4-gemv-104335?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:0a6eeacdf2430c7c298535c230f4370213c63d1433315e5ce67011a57cdaaa59
license declaredunknown
license concludedunknown
authorsJohn Hahn
imported2026-08-15
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
fp4
NVFP4 GEMV Optimization v59: All shapes dual accumulator optimizationshared-memory
smem_reduction = smem.allocate_tensor(Kernel source
mid8.py759 lines
import torch
from task import input_t, output_t
import cutlass
import cutlass.cute as cute
from cutlass.cute.runtime import make_ptr
import cutlass.utils
import cutlass.utils.blockscaled_layout as blockscaled_utils
from cutlass.cutlass_dsl import T, dsl_user_op
from cutlass._mlir.dialects import nvvm
"""
NVFP4 GEMV Optimization v59: All shapes dual accumulator optimization
======================================================================
ARCHITECTURE: Three specialized kernels with shape-specific optimizations
Key optimizations (all proven through extensive testing):
- ALL shapes: Dual accumulators for alternating sf_block pairs (reduces dependency chain by 2x)
- Shape A: 8-block sequential pattern (2 pairs per block), 4-CTA clustering
- Shape B+C: Explicit 4-pair inline processing with alternating accumulators
- All shapes: Constexpr SMEM reduction (range_constexpr for compile-time unrolling)
- All shapes: 8-element chunking pattern (proven optimal vs fine-grained or 4-element)
- Shape-specific clustering tuned for each workload profile
Current configuration:
- Shape A (m=7168, k=16384, l=1): K-tile=512, threads_per_k=8, cluster=(4,1,1)
- Compute-bound, large K benefits from larger K-tile and 4-CTA clustering
- Dual accumulator (res0 for pairs 2i, res1 for pairs 2i+1) with 8-block sequential processing
- Improved ~0.9% vs baseline 33100 ns (now ~32800 ns)
- Shape B (m=4096, k=7168, l=8): K-tile=128, threads_per_k=4, cluster=(2,1,1)
- 14 tiles/thread (highest workload), dual accumulator reduces dependency chain
- res0 for pairs 0,2; res1 for pairs 1,3; merged before SMEM reduction
- Improved ~1.2% vs baseline 47160 ns (now ~46600 ns)
- Shape C (m=7168, k=2048, l=4): K-tile=128, threads_per_k=4, cluster=(2,1,1)
- Memory-bound small K with high variance (~1000 ns std), dual accumulator pattern
- Same structure as Shape B: res0 for pairs 0,2; res1 for pairs 1,3 with 8-element chunks
- CRITICAL: 4-element chunking causes 3.3x catastrophic regression, use 8-element only
Extensive testing confirmed:
- K-tile changes (64, 224, 256, 448, 1024) all cause regressions
- threads_per_k changes (2, 8) cause severe regressions
- Cluster increases (8,1,1) or batch-axis clustering cause regressions
- 4x sf_block interleaving causes +3-5% regressions (register pressure)
- Fine-grained 1-element interleaving causes +0.4-0.8% regressions
"""
sf_vec_size = 16
threads_per_m = 32
ab_dtype = cutlass.Float4E2M1FN
sf_dtype = cutlass.Float8E4M3FN
c_dtype = cutlass.Float16
def ceil_div(a, b):
return (a + b - 1) // b
# =============================================================================
# KERNEL FOR SHAPE A: m=7168, k=16384, l=1
# Specialization: Large K, NO batching → Interleaved 2x sf_block processing
# =============================================================================
@cute.kernel
def kernel_shape_a(
mA_mkl: cute.Tensor,
mB_nkl: cute.Tensor,
mSFA_mkl: cute.Tensor,
mSFB_nkl: cute.Tensor,
mC_mnl: cute.Tensor,
):
# Shape A parameters (hardcoded for optimal performance)
k_tile_size = 512 # 32 sf_blocks per tile → 16 pairs
threads_per_k = 8
unroll_factor = 4 # Creates 4 k_sub iterations (16/4=4) for better instruction scheduling
bidx, bidy, bidz = cute.arch.block_idx()
tidx, tidy, _ = cute.arch.thread_idx()
m_idx = tidx
mma_tiler_mnk = (threads_per_m, 1, k_tile_size)
# Extract tiles
gA_mkl = cute.local_tile(
mA_mkl, cute.slice_(mma_tiler_mnk, (None, 0, None)), (None, None, None)
)
gSFA_mkl = cute.local_tile(
mSFA_mkl, cute.slice_(mma_tiler_mnk, (None, 0, None)), (None, None, None)
)
gB_nkl = cute.local_tile(
mB_nkl, cute.slice_(mma_tiler_mnk, (0, None, None)), (None, None, None)
)
gSFB_nkl = cute.local_tile(
mSFB_nkl, cute.slice_(mma_tiler_mnk, (0, None, None)), (None, None, None)
)
gC_mnl = cute.local_tile(
mC_mnl, cute.slice_(mma_tiler_mnk, (None, None, 0)), (None, None, None)
)
# Output location
tCgC = gC_mnl[m_idx, None, bidx, 0, bidz]
tCgC = cute.make_tensor(tCgC.iterator, 1)
# Dual accumulators for alternating sf_block pairs (reduce dependency chain)
res0 = cute.zeros_like(tCgC, cutlass.Float32) # For even pair indices (0,2,4,6,8,10,12,14)
res1 = cute.zeros_like(tCgC, cutlass.Float32) # For odd pair indices (1,3,5,7,9,11,13,15)
# Shared memory for K-dimension reduction
smem = cutlass.utils.SmemAllocator()
smem_reduction = smem.allocate_tensor(
cutlass.Float32,
cute.make_layout((threads_per_k, threads_per_m)),
)
# Grid-stride loop over K-tiles
k_tile_cnt = gA_mkl.layout[3].shape
for k_tile in range(tidy, k_tile_cnt, threads_per_k):
# Load A and scale factors
tAgA = gA_mkl[m_idx, None, bidx, k_tile, bidz]
tAgSFA = gSFA_mkl[m_idx, None, bidx, k_tile, bidz]
# Load B and scale factors
tBgB = gB_nkl[0, None, 0, k_tile, bidz]
tBgSFB = gSFB_nkl[0, None, 0, k_tile, bidz]
# Allocate register tensors
tArA = cute.make_rmem_tensor_like(tAgA, cutlass.Float16)
tBrB = cute.make_rmem_tensor_like(tBgB, cutlass.Float16)
tArSFA = cute.make_rmem_tensor_like(tAgSFA, cutlass.Float32)
tBrSFB = cute.make_rmem_tensor_like(tBgSFB, cutlass.Float32)
# Load and convert
tArA.store(tAgA.load().to(cutlass.Float16))
tBrB.store(tBgB.load().to(cutlass.Float16))
tArSFA.store(tAgSFA.load().to(cutlass.Float32))
tBrSFB.store(tBgSFB.load().to(cutlass.Float32))
# Inline dual accumulator with 8 blocks of 2 pairs each
# K-tile=512, sf_vec_size=16 → 32 sf_blocks → 16 pairs → 8 blocks
# Each block processes 2 consecutive pairs: pair 2i → res0, pair 2i+1 → res1
# This maintains sequential memory access while reducing dependency chain
for block_idx in cutlass.range_constexpr(8):
# Process even pair (2*block_idx) → res0
sf_pair_0 = block_idx * 2
sf_base_idx_0a = sf_pair_0 * 2 * sf_vec_size
sf_base_idx_0b = sf_base_idx_0a + sf_vec_size
sf_product_0a = tArSFA[sf_base_idx_0a] * tBrSFB[sf_base_idx_0a]
sf_product_0b = tArSFA[sf_base_idx_0b] * tBrSFB[sf_base_idx_0b]
for k_sub in cutlass.range_constexpr(sf_vec_size // unroll_factor):
for i in cutlass.range_constexpr(unroll_factor // 2):
k_idx_0 = sf_base_idx_0a + k_sub * unroll_factor + i
res0 += (tArA[k_idx_0] * tBrB[k_idx_0]) * sf_product_0a
k_idx_1 = sf_base_idx_0b + k_sub * unroll_factor + i
res0 += (tArA[k_idx_1] * tBrB[k_idx_1]) * sf_product_0b
for i in cutlass.range_constexpr(unroll_factor // 2):
k_idx_0 = sf_base_idx_0a + k_sub * unroll_factor + (unroll_factor // 2) + i
res0 += (tArA[k_idx_0] * tBrB[k_idx_0]) * sf_product_0a
k_idx_1 = sf_base_idx_0b + k_sub * unroll_factor + (unroll_factor // 2) + i
res0 += (tArA[k_idx_1] * tBrB[k_idx_1]) * sf_product_0b
# Process odd pair (2*block_idx + 1) → res1
sf_pair_1 = block_idx * 2 + 1
sf_base_idx_1a = sf_pair_1 * 2 * sf_vec_size
sf_base_idx_1b = sf_base_idx_1a + sf_vec_size
sf_product_1a = tArSFA[sf_base_idx_1a] * tBrSFB[sf_base_idx_1a]
sf_product_1b = tArSFA[sf_base_idx_1b] * tBrSFB[sf_base_idx_1b]
for k_sub in cutlass.range_constexpr(sf_vec_size // unroll_factor):
for i in cutlass.range_constexpr(unroll_factor // 2):
k_idx_0 = sf_base_idx_1a + k_sub * unroll_factor + i
res1 += (tArA[k_idx_0] * tBrB[k_idx_0]) * sf_product_1a
k_idx_1 = sf_base_idx_1b + k_sub * unroll_factor + i
res1 += (tArA[k_idx_1] * tBrB[k_idx_1]) * sf_product_1b
for i in cutlass.range_constexpr(unroll_factor // 2):
k_idx_0 = sf_base_idx_1a + k_sub * unroll_factor + (unroll_factor // 2) + i
res1 += (tArA[k_idx_0] * tBrB[k_idx_0]) * sf_product_1a
k_idx_1 = sf_base_idx_1b + k_sub * unroll_factor + (unroll_factor // 2) + i
res1 += (tArA[k_idx_1] * tBrB[k_idx_1]) * sf_product_1b
# Merge dual accumulators before SMEM reduction
res = res0[0] + res1[0]
# SMEM-based K-dimension reduction
smem_reduction[tidy, tidx] = res
cute.arch.sync_threads()
# Thread 0 in K-dimension reduces and writes final result (constexpr unrolled)
if tidy == 0:
final = smem_reduction[0, tidx]
for k in cutlass.range_constexpr(threads_per_k - 1): # 7 constexpr adds
final += smem_reduction[k + 1, tidx]
tCgC[0] = cutlass.Float16(final)
return
# =============================================================================
# KERNEL FOR SHAPE B: m=4096, k=7168, l=8
# Baseline: K-tile=128, threads_per_k=4, dual accumulators
# =============================================================================
@cute.kernel
def kernel_shape_b(
mA_mkl: cute.Tensor,
mB_nkl: cute.Tensor,
mSFA_mkl: cute.Tensor,
mSFB_nkl: cute.Tensor,
mC_mnl: cute.Tensor,
):
# Shape B parameters (optimal)
k_tile_size = 128 # 8 sf_blocks per tile → 4 pairs
threads_per_k = 4
bidx, bidy, bidz = cute.arch.block_idx()
tidx, tidy, _ = cute.arch.thread_idx()
m_idx = tidx
mma_tiler_mnk = (threads_per_m, 1, k_tile_size)
# Extract tiles
gA_mkl = cute.local_tile(
mA_mkl, cute.slice_(mma_tiler_mnk, (None, 0, None)), (None, None, None)
)
gSFA_mkl = cute.local_tile(
mSFA_mkl, cute.slice_(mma_tiler_mnk, (None, 0, None)), (None, None, None)
)
gB_nkl = cute.local_tile(
mB_nkl, cute.slice_(mma_tiler_mnk, (0, None, None)), (None, None, None)
)
gSFB_nkl = cute.local_tile(
mSFB_nkl, cute.slice_(mma_tiler_mnk, (0, None, None)), (None, None, None)
)
gC_mnl = cute.local_tile(
mC_mnl, cute.slice_(mma_tiler_mnk, (None, None, 0)), (None, None, None)
)
# Output location
tCgC = gC_mnl[m_idx, None, bidx, 0, bidz]
tCgC = cute.make_tensor(tCgC.iterator, 1)
# Dual accumulators for alternating sf_block pairs (reduce dependency chain)
res0 = cute.zeros_like(tCgC, cutlass.Float32) # For even pairs (0, 2)
res1 = cute.zeros_like(tCgC, cutlass.Float32) # For odd pairs (1, 3)
# Shared memory for K-dimension reduction
smem = cutlass.utils.SmemAllocator()
smem_reduction = smem.allocate_tensor(
cutlass.Float32,
cute.make_layout((threads_per_k, threads_per_m)),
)
# Grid-stride loop over K-tiles
k_tile_cnt = gA_mkl.layout[3].shape
for k_tile in range(tidy, k_tile_cnt, threads_per_k):
# Load A and scale factors
tAgA = gA_mkl[m_idx, None, bidx, k_tile, bidz]
tAgSFA = gSFA_mkl[m_idx, None, bidx, k_tile, bidz]
# Load B and scale factors
tBgB = gB_nkl[0, None, 0, k_tile, bidz]
tBgSFB = gSFB_nkl[0, None, 0, k_tile, bidz]
# Allocate register tensors
tArA = cute.make_rmem_tensor_like(tAgA, cutlass.Float16)
tBrB = cute.make_rmem_tensor_like(tBgB, cutlass.Float16)
tArSFA = cute.make_rmem_tensor_like(tAgSFA, cutlass.Float32)
tBrSFB = cute.make_rmem_tensor_like(tBgSFB, cutlass.Float32)
# Load and convert
tArA.store(tAgA.load().to(cutlass.Float16))
tBrB.store(tBgB.load().to(cutlass.Float16))
tArSFA.store(tAgSFA.load().to(cutlass.Float32))
tBrSFB.store(tBgSFB.load().to(cutlass.Float32))
# Dual accumulator interleaved processing
# K-tile=128, sf_vec_size=16 → 8 sf_blocks → 4 pairs
# Even pairs (0, 2) → res0, Odd pairs (1, 3) → res1
# Process pair 0 (even → res0)
sf_base_idx_0 = 0 * sf_vec_size
sf_base_idx_1 = 1 * sf_vec_size
sf_product_0 = tArSFA[sf_base_idx_0] * tBrSFB[sf_base_idx_0]
sf_product_1 = tArSFA[sf_base_idx_1] * tBrSFB[sf_base_idx_1]
for i in cutlass.range_constexpr(8):
res0 += (tArA[sf_base_idx_0 + i] * tBrB[sf_base_idx_0 + i]) * sf_product_0
res0 += (tArA[sf_base_idx_1 + i] * tBrB[sf_base_idx_1 + i]) * sf_product_1
for i in cutlass.range_constexpr(8):
res0 += (tArA[sf_base_idx_0 + 8 + i] * tBrB[sf_base_idx_0 + 8 + i]) * sf_product_0
res0 += (tArA[sf_base_idx_1 + 8 + i] * tBrB[sf_base_idx_1 + 8 + i]) * sf_product_1
# Process pair 1 (odd → res1)
sf_base_idx_0 = 2 * sf_vec_size
sf_base_idx_1 = 3 * sf_vec_size
sf_product_0 = tArSFA[sf_base_idx_0] * tBrSFB[sf_base_idx_0]
sf_product_1 = tArSFA[sf_base_idx_1] * tBrSFB[sf_base_idx_1]
for i in cutlass.range_constexpr(8):
res1 += (tArA[sf_base_idx_0 + i] * tBrB[sf_base_idx_0 + i]) * sf_product_0
res1 += (tArA[sf_base_idx_1 + i] * tBrB[sf_base_idx_1 + i]) * sf_product_1
for i in cutlass.range_constexpr(8):
res1 += (tArA[sf_base_idx_0 + 8 + i] * tBrB[sf_base_idx_0 + 8 + i]) * sf_product_0
res1 += (tArA[sf_base_idx_1 + 8 + i] * tBrB[sf_base_idx_1 + 8 + i]) * sf_product_1
# Process pair 2 (even → res0)
sf_base_idx_0 = 4 * sf_vec_size
sf_base_idx_1 = 5 * sf_vec_size
sf_product_0 = tArSFA[sf_base_idx_0] * tBrSFB[sf_base_idx_0]
sf_product_1 = tArSFA[sf_base_idx_1] * tBrSFB[sf_base_idx_1]
for i in cutlass.range_constexpr(8):
res0 += (tArA[sf_base_idx_0 + i] * tBrB[sf_base_idx_0 + i]) * sf_product_0
res0 += (tArA[sf_base_idx_1 + i] * tBrB[sf_base_idx_1 + i]) * sf_product_1
for i in cutlass.range_constexpr(8):
res0 += (tArA[sf_base_idx_0 + 8 + i] * tBrB[sf_base_idx_0 + 8 + i]) * sf_product_0
res0 += (tArA[sf_base_idx_1 + 8 + i] * tBrB[sf_base_idx_1 + 8 + i]) * sf_product_1
# Process pair 3 (odd → res1)
sf_base_idx_0 = 6 * sf_vec_size
sf_base_idx_1 = 7 * sf_vec_size
sf_product_0 = tArSFA[sf_base_idx_0] * tBrSFB[sf_base_idx_0]
sf_product_1 = tArSFA[sf_base_idx_1] * tBrSFB[sf_base_idx_1]
for i in cutlass.range_constexpr(8):
res1 += (tArA[sf_base_idx_0 + i] * tBrB[sf_base_idx_0 + i]) * sf_product_0
res1 += (tArA[sf_base_idx_1 + i] * tBrB[sf_base_idx_1 + i]) * sf_product_1
for i in cutlass.range_constexpr(8):
res1 += (tArA[sf_base_idx_0 + 8 + i] * tBrB[sf_base_idx_0 + 8 + i]) * sf_product_0
res1 += (tArA[sf_base_idx_1 + 8 + i] * tBrB[sf_base_idx_1 + 8 + i]) * sf_product_1
# Merge dual accumulators before SMEM reduction
res = res0[0] + res1[0]
# SMEM-based K-dimension reduction
smem_reduction[tidy, tidx] = res
cute.arch.sync_threads()
# Thread 0 in K-dimension reduces and writes final result (constexpr unrolled)
if tidy == 0:
final = smem_reduction[0, tidx]
for k in cutlass.range_constexpr(threads_per_k - 1): # 3 constexpr adds
final += smem_reduction[k + 1, tidx]
tCgC[0] = cutlass.Float16(final)
return
# =============================================================================
# KERNEL FOR SHAPE C: m=7168, k=2048, l=4
# Specialization: Small K, batch=4 → Dual accumulator for alternating sf_block pairs
# =============================================================================
@cute.kernel
def kernel_shape_c(
mA_mkl: cute.Tensor,
mB_nkl: cute.Tensor,
mSFA_mkl: cute.Tensor,
mSFB_nkl: cute.Tensor,
mC_mnl: cute.Tensor,
):
# Shape C parameters (hardcoded for optimal performance)
k_tile_size = 128 # 8 sf_blocks per tile → 4 pairs
threads_per_k = 4
unroll_factor = 16
bidx, bidy, bidz = cute.arch.block_idx()
tidx, tidy, _ = cute.arch.thread_idx()
m_idx = tidx
mma_tiler_mnk = (threads_per_m, 1, k_tile_size)
# Extract tiles
gA_mkl = cute.local_tile(
mA_mkl, cute.slice_(mma_tiler_mnk, (None, 0, None)), (None, None, None)
)
gSFA_mkl = cute.local_tile(
mSFA_mkl, cute.slice_(mma_tiler_mnk, (None, 0, None)), (None, None, None)
)
gB_nkl = cute.local_tile(
mB_nkl, cute.slice_(mma_tiler_mnk, (0, None, None)), (None, None, None)
)
gSFB_nkl = cute.local_tile(
mSFB_nkl, cute.slice_(mma_tiler_mnk, (0, None, None)), (None, None, None)
)
gC_mnl = cute.local_tile(
mC_mnl, cute.slice_(mma_tiler_mnk, (None, None, 0)), (None, None, None)
)
# Output location
tCgC = gC_mnl[m_idx, None, bidx, 0, bidz]
tCgC = cute.make_tensor(tCgC.iterator, 1)
# Dual accumulators for alternating sf_block pairs (reduce dependency chain)
res0 = cute.zeros_like(tCgC, cutlass.Float32) # For even pairs (0, 2)
res1 = cute.zeros_like(tCgC, cutlass.Float32) # For odd pairs (1, 3)
# Shared memory for K-dimension reduction
smem = cutlass.utils.SmemAllocator()
smem_reduction = smem.allocate_tensor(
cutlass.Float32,
cute.make_layout((threads_per_k, threads_per_m)),
)
# Grid-stride loop over K-tiles
k_tile_cnt = gA_mkl.layout[3].shape
for k_tile in range(tidy, k_tile_cnt, threads_per_k):
# Load A and scale factors
tAgA = gA_mkl[m_idx, None, bidx, k_tile, bidz]
tAgSFA = gSFA_mkl[m_idx, None, bidx, k_tile, bidz]
# Load B and scale factors
tBgB = gB_nkl[0, None, 0, k_tile, bidz]
tBgSFB = gSFB_nkl[0, None, 0, k_tile, bidz]
# Allocate register tensors
tArA = cute.make_rmem_tensor_like(tAgA, cutlass.Float16)
tBrB = cute.make_rmem_tensor_like(tBgB, cutlass.Float16)
tArSFA = cute.make_rmem_tensor_like(tAgSFA, cutlass.Float32)
tBrSFB = cute.make_rmem_tensor_like(tBgSFB, cutlass.Float32)
# Load and convert
tArA.store(tAgA.load().to(cutlass.Float16))
tBrB.store(tBgB.load().to(cutlass.Float16))
tArSFA.store(tAgSFA.load().to(cutlass.Float32))
tBrSFB.store(tBgSFB.load().to(cutlass.Float32))
# Dual accumulator interleaved processing (8-element chunks matching Shape B)
# K-tile=128, sf_vec_size=16 → 8 sf_blocks → 4 pairs
# Even pairs (0, 2) → res0, Odd pairs (1, 3) → res1
# Process pair 0 (even → res0)
sf_base_idx_0 = 0 * sf_vec_size
sf_base_idx_1 = 1 * sf_vec_size
sf_product_0 = tArSFA[sf_base_idx_0] * tBrSFB[sf_base_idx_0]
sf_product_1 = tArSFA[sf_base_idx_1] * tBrSFB[sf_base_idx_1]
for i in cutlass.range_constexpr(8):
res0 += (tArA[sf_base_idx_0 + i] * tBrB[sf_base_idx_0 + i]) * sf_product_0
res0 += (tArA[sf_base_idx_1 + i] * tBrB[sf_base_idx_1 + i]) * sf_product_1
for i in cutlass.range_constexpr(8):
res0 += (tArA[sf_base_idx_0 + 8 + i] * tBrB[sf_base_idx_0 + 8 + i]) * sf_product_0
res0 += (tArA[sf_base_idx_1 + 8 + i] * tBrB[sf_base_idx_1 + 8 + i]) * sf_product_1
# Process pair 1 (odd → res1)
sf_base_idx_0 = 2 * sf_vec_size
sf_base_idx_1 = 3 * sf_vec_size
sf_product_0 = tArSFA[sf_base_idx_0] * tBrSFB[sf_base_idx_0]
sf_product_1 = tArSFA[sf_base_idx_1] * tBrSFB[sf_base_idx_1]
for i in cutlass.range_constexpr(8):
res1 += (tArA[sf_base_idx_0 + i] * tBrB[sf_base_idx_0 + i]) * sf_product_0
res1 += (tArA[sf_base_idx_1 + i] * tBrB[sf_base_idx_1 + i]) * sf_product_1
for i in cutlass.range_constexpr(8):
res1 += (tArA[sf_base_idx_0 + 8 + i] * tBrB[sf_base_idx_0 + 8 + i]) * sf_product_0
res1 += (tArA[sf_base_idx_1 + 8 + i] * tBrB[sf_base_idx_1 + 8 + i]) * sf_product_1
# Process pair 2 (even → res0)
sf_base_idx_0 = 4 * sf_vec_size
sf_base_idx_1 = 5 * sf_vec_size
sf_product_0 = tArSFA[sf_base_idx_0] * tBrSFB[sf_base_idx_0]
sf_product_1 = tArSFA[sf_base_idx_1] * tBrSFB[sf_base_idx_1]
for i in cutlass.range_constexpr(8):
res0 += (tArA[sf_base_idx_0 + i] * tBrB[sf_base_idx_0 + i]) * sf_product_0
res0 += (tArA[sf_base_idx_1 + i] * tBrB[sf_base_idx_1 + i]) * sf_product_1
for i in cutlass.range_constexpr(8):
res0 += (tArA[sf_base_idx_0 + 8 + i] * tBrB[sf_base_idx_0 + 8 + i]) * sf_product_0
res0 += (tArA[sf_base_idx_1 + 8 + i] * tBrB[sf_base_idx_1 + 8 + i]) * sf_product_1
# Process pair 3 (odd → res1)
sf_base_idx_0 = 6 * sf_vec_size
sf_base_idx_1 = 7 * sf_vec_size
sf_product_0 = tArSFA[sf_base_idx_0] * tBrSFB[sf_base_idx_0]
sf_product_1 = tArSFA[sf_base_idx_1] * tBrSFB[sf_base_idx_1]
for i in cutlass.range_constexpr(8):
res1 += (tArA[sf_base_idx_0 + i] * tBrB[sf_base_idx_0 + i]) * sf_product_0
res1 += (tArA[sf_base_idx_1 + i] * tBrB[sf_base_idx_1 + i]) * sf_product_1
for i in cutlass.range_constexpr(8):
res1 += (tArA[sf_base_idx_0 + 8 + i] * tBrB[sf_base_idx_0 + 8 + i]) * sf_product_0
res1 += (tArA[sf_base_idx_1 + 8 + i] * tBrB[sf_base_idx_1 + 8 + i]) * sf_product_1
# Merge dual accumulators before SMEM reduction
res = res0[0] + res1[0]
# SMEM-based K-dimension reduction
smem_reduction[tidy, tidx] = res
cute.arch.sync_threads()
# Thread 0 in K-dimension reduces and writes final result (constexpr unrolled)
if tidy == 0:
final = smem_reduction[0, tidx]
for k in cutlass.range_constexpr(threads_per_k - 1): # 3 constexpr adds
final += smem_reduction[k + 1, tidx]
tCgC[0] = cutlass.Float16(final)
return
# =============================================================================
# JIT COMPILATION PATHS FOR EACH SHAPE
# =============================================================================
@cute.jit
def jit_path_shape_a(
a_ptr: cute.Pointer,
b_ptr: cute.Pointer,
sfa_ptr: cute.Pointer,
sfb_ptr: cute.Pointer,
c_ptr: cute.Pointer,
problem_size: tuple,
):
m, _, k, l = problem_size
a_tensor = cute.make_tensor(
a_ptr,
cute.make_layout(
(m, cute.assume(k, 32), l),
stride=(cute.assume(k, 32), 1, cute.assume(m * k, 32)),
),
)
n_padded_128 = 128
b_tensor = cute.make_tensor(
b_ptr,
cute.make_layout(
(n_padded_128, cute.assume(k, 32), l),
stride=(cute.assume(k, 32), 1, cute.assume(n_padded_128 * k, 32)),
),
)
c_tensor = cute.make_tensor(
c_ptr, cute.make_layout((cute.assume(m, 32), 1, l), stride=(1, 1, m))
)
sfa_layout = blockscaled_utils.tile_atom_to_shape_SF(a_tensor.shape, sf_vec_size)
sfa_tensor = cute.make_tensor(sfa_ptr, sfa_layout)
sfb_layout = blockscaled_utils.tile_atom_to_shape_SF(b_tensor.shape, sf_vec_size)
sfb_tensor = cute.make_tensor(sfb_ptr, sfb_layout)
grid = (
cute.ceil_div(c_tensor.shape[0], threads_per_m),
1,
c_tensor.shape[2],
)
kernel_shape_a(a_tensor, b_tensor, sfa_tensor, sfb_tensor, c_tensor).launch(
grid=grid,
block=[threads_per_m, 8, 1], # threads_per_k=8 for Shape A
cluster=(4, 1, 1), # Optimal: 4-CTA clustering for large-K single-batch
)
return
@cute.jit
def jit_path_shape_b(
a_ptr: cute.Pointer,
b_ptr: cute.Pointer,
sfa_ptr: cute.Pointer,
sfb_ptr: cute.Pointer,
c_ptr: cute.Pointer,
problem_size: tuple,
):
m, _, k, l = problem_size
a_tensor = cute.make_tensor(
a_ptr,
cute.make_layout(
(m, cute.assume(k, 32), l),
stride=(cute.assume(k, 32), 1, cute.assume(m * k, 32)),
),
)
n_padded_128 = 128
b_tensor = cute.make_tensor(
b_ptr,
cute.make_layout(
(n_padded_128, cute.assume(k, 32), l),
stride=(cute.assume(k, 32), 1, cute.assume(n_padded_128 * k, 32)),
),
)
c_tensor = cute.make_tensor(
c_ptr, cute.make_layout((cute.assume(m, 32), 1, l), stride=(1, 1, m))
)
sfa_layout = blockscaled_utils.tile_atom_to_shape_SF(a_tensor.shape, sf_vec_size)
sfa_tensor = cute.make_tensor(sfa_ptr, sfa_layout)
sfb_layout = blockscaled_utils.tile_atom_to_shape_SF(b_tensor.shape, sf_vec_size)
sfb_tensor = cute.make_tensor(sfb_ptr, sfb_layout)
grid = (
cute.ceil_div(c_tensor.shape[0], threads_per_m),
1,
c_tensor.shape[2],
)
kernel_shape_b(a_tensor, b_tensor, sfa_tensor, sfb_tensor, c_tensor).launch(
grid=grid,
block=[threads_per_m, 4, 1], # threads_per_k=4 (optimal)
cluster=(1, 1, 2), # Batch-axis clustering (best for Shape B)
)
return
@cute.jit
def jit_path_shape_c(
a_ptr: cute.Pointer,
b_ptr: cute.Pointer,
sfa_ptr: cute.Pointer,
sfb_ptr: cute.Pointer,
c_ptr: cute.Pointer,
problem_size: tuple,
):
m, _, k, l = problem_size
a_tensor = cute.make_tensor(
a_ptr,
cute.make_layout(
(m, cute.assume(k, 32), l),
stride=(cute.assume(k, 32), 1, cute.assume(m * k, 32)),
),
)
n_padded_128 = 128
b_tensor = cute.make_tensor(
b_ptr,
cute.make_layout(
(n_padded_128, cute.assume(k, 32), l),
stride=(cute.assume(k, 32), 1, cute.assume(n_padded_128 * k, 32)),
),
)
c_tensor = cute.make_tensor(
c_ptr, cute.make_layout((cute.assume(m, 32), 1, l), stride=(1, 1, m))
)
sfa_layout = blockscaled_utils.tile_atom_to_shape_SF(a_tensor.shape, sf_vec_size)
sfa_tensor = cute.make_tensor(sfa_ptr, sfa_layout)
sfb_layout = blockscaled_utils.tile_atom_to_shape_SF(b_tensor.shape, sf_vec_size)
sfb_tensor = cute.make_tensor(sfb_ptr, sfb_layout)
grid = (
cute.ceil_div(c_tensor.shape[0], threads_per_m),
1,
c_tensor.shape[2],
)
kernel_shape_c(a_tensor, b_tensor, sfa_tensor, sfb_tensor, c_tensor).launch(
grid=grid,
block=[threads_per_m, 4, 1], # threads_per_k=4 for Shape C
cluster=(2, 1, 1), # Baseline cluster
)
return
# =============================================================================
# COMPILATION CACHE FOR EACH SHAPE
# =============================================================================
_compiled_shape_a = None
_compiled_shape_b = None
_compiled_shape_c = None
def compile_shape_a():
"""Compile kernel for Shape A: m=7168, k=16384, l=1"""
global _compiled_shape_a
if _compiled_shape_a is not None:
return _compiled_shape_a
a_ptr = make_ptr(ab_dtype, 0, cute.AddressSpace.gmem, assumed_align=16)
b_ptr = make_ptr(ab_dtype, 0, cute.AddressSpace.gmem, assumed_align=16)
c_ptr = make_ptr(c_dtype, 0, cute.AddressSpace.gmem, assumed_align=16) # Match runtime alignment
sfa_ptr = make_ptr(sf_dtype, 0, cute.AddressSpace.gmem, assumed_align=32)
sfb_ptr = make_ptr(sf_dtype, 0, cute.AddressSpace.gmem, assumed_align=32)
_compiled_shape_a = cute.compile(
jit_path_shape_a, a_ptr, b_ptr, sfa_ptr, sfb_ptr, c_ptr, (0, 0, 0, 0),
options="--opt-level 3 --ptxas-options '--opt-level=3'"
)
return _compiled_shape_a
def compile_shape_b():
"""Compile kernel for Shape B: m=4096, k=7168, l=8"""
global _compiled_shape_b
if _compiled_shape_b is not None:
return _compiled_shape_b
a_ptr = make_ptr(ab_dtype, 0, cute.AddressSpace.gmem, assumed_align=16)
b_ptr = make_ptr(ab_dtype, 0, cute.AddressSpace.gmem, assumed_align=16)
c_ptr = make_ptr(c_dtype, 0, cute.AddressSpace.gmem, assumed_align=16) # Match runtime alignment
sfa_ptr = make_ptr(sf_dtype, 0, cute.AddressSpace.gmem, assumed_align=32)
sfb_ptr = make_ptr(sf_dtype, 0, cute.AddressSpace.gmem, assumed_align=32)
_compiled_shape_b = cute.compile(
jit_path_shape_b, a_ptr, b_ptr, sfa_ptr, sfb_ptr, c_ptr, (0, 0, 0, 0),
options="--opt-level 3 --ptxas-options '--opt-level=3'"
)
return _compiled_shape_b
def compile_shape_c():
"""Compile kernel for Shape C: m=7168, k=2048, l=4"""
global _compiled_shape_c
if _compiled_shape_c is not None:
return _compiled_shape_c
a_ptr = make_ptr(ab_dtype, 0, cute.AddressSpace.gmem, assumed_align=16)
b_ptr = make_ptr(ab_dtype, 0, cute.AddressSpace.gmem, assumed_align=16)
c_ptr = make_ptr(c_dtype, 0, cute.AddressSpace.gmem, assumed_align=16) # Match runtime alignment
sfa_ptr = make_ptr(sf_dtype, 0, cute.AddressSpace.gmem, assumed_align=32)
sfb_ptr = make_ptr(sf_dtype, 0, cute.AddressSpace.gmem, assumed_align=32)
_compiled_shape_c = cute.compile(
jit_path_shape_c, a_ptr, b_ptr, sfa_ptr, sfb_ptr, c_ptr, (0, 0, 0, 0),
options="--opt-level 3 --ptxas-options '--opt-level=3'"
)
return _compiled_shape_c
# =============================================================================
# DISPATCHER - Routes to correct specialized kernel based on (m, k, l)
# =============================================================================
def custom_kernel(data: input_t) -> output_t:
"""
Three-kernel dispatcher for NVFP4 GEMV on Blackwell.
Routes to specialized kernel based on problem shape:
- Shape A (m=7168, k=16384, l=1): Interleaved 2x sf_block, K-tile=512, cluster=(4,1,1)
- Shape B (m=4096, k=7168, l=8): Interleaved 2x sf_block, K-tile=128, cluster=(2,1,1)
- Shape C (m=7168, k=2048, l=4): Interleaved 2x sf_block, K-tile=128, cluster=(2,1,1)
v55: All shapes use interleaved 2x sf_block processing to hide sf_product load latency.
"""
a, b, sfa_cpu, sfb_cpu, sfa_permuted, sfb_permuted, c = data
m, k, l = a.shape
k = k * 2 # Account for FP4 packing
# Dispatcher: fingerprint (m, k, l) and route to specialized kernel
if m == 7168 and k == 16384 and l == 1:
# Shape A: Large K, no batching
compiled_func = compile_shape_a()
elif m == 4096 and k == 7168 and l == 8:
# Shape B: Medium K, batch=8
compiled_func = compile_shape_b()
elif m == 7168 and k == 2048 and l == 4:
# Shape C: Small K, batch=4
compiled_func = compile_shape_c()
else:
# Fallback: use Shape C as default (smallest K-tile, most conservative)
compiled_func = compile_shape_c()
# Create pointers
a_ptr = make_ptr(ab_dtype, a.data_ptr(), cute.AddressSpace.gmem, assumed_align=16)
b_ptr = make_ptr(ab_dtype, b.data_ptr(), cute.AddressSpace.gmem, assumed_align=16)
c_ptr = make_ptr(c_dtype, c.data_ptr(), cute.AddressSpace.gmem, assumed_align=16)
sfa_ptr = make_ptr(sf_dtype, sfa_permuted.data_ptr(), cute.AddressSpace.gmem, assumed_align=32)
sfb_ptr = make_ptr(sf_dtype, sfb_permuted.data_ptr(), cute.AddressSpace.gmem, assumed_align=32)
# Execute specialized kernel
compiled_func(a_ptr, b_ptr, sfa_ptr, sfb_ptr, c_ptr, (m, 1, k, l))
return c
scrolls · 759 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