submission 96007
RIM#0013 · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 1053 lines, June 9 Researcher Reciprocity License v1.0.
new.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-nvfp4-gemv-96007?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:d8ac4dac84ebf7deb658c3b44d0e86e63737c364effe1d9f1960c2be6613eb0b
license declaredunknown
license concludedunknown
authorsRIM#0013
imported2026-08-26
Kernel source
new.py1053 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.blockscaled_layout as blockscaled_utils
# # Blackwell-optimized kernel configuration
# mma_tiler_mnk = (256, 1, 128) # Larger tiles for Blackwell's increased SM count and register file
# ab_dtype = cutlass.Float4E2M1FN
# sf_dtype = cutlass.Float8E4M3FN
# c_dtype = cutlass.Float16
# sf_vec_size = 16
# threads_per_cta = 256 # Increased thread count for better occupancy on Blackwell
# def ceil_div(a, b):
# return (a + b - 1) // b
# @cute.kernel
# def kernel(
# mA_mkl: cute.Tensor,
# mB_nkl: cute.Tensor,
# mSFA_mkl: cute.Tensor,
# mSFB_nkl: cute.Tensor,
# mC_mnl: cute.Tensor,
# ):
# bidx, bidy, bidz = cute.arch.block_idx()
# tidx, _, _ = cute.arch.thread_idx()
# # Local tile extraction with larger tile sizes
# 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)
# )
# tCgC = gC_mnl[tidx, None, bidx, bidy, bidz]
# tCgC = cute.make_tensor(tCgC.iterator, 1)
# # Use FP32 accumulation for better precision on Blackwell
# res = cute.zeros_like(tCgC, cutlass.Float32)
# k_tile_cnt = gA_mkl.layout[3].shape
# # Unroll factor optimized for Blackwell's instruction cache and pipeline depth
# unroll_factor = 4
# k_tile_unrolled = (k_tile_cnt // unroll_factor) * unroll_factor
# # Main unrolled loop for better ILP on Blackwell
# for k_tile in cutlass.range_constexpr(0, k_tile_unrolled, unroll_factor):
# # Process 4 k-tiles per iteration for better instruction-level parallelism
# for unroll_idx in cutlass.range_constexpr(unroll_factor):
# k_idx = k_tile + unroll_idx
# tAgA = gA_mkl[tidx, None, bidx, k_idx, bidz]
# tBgB = gB_nkl[0, None, bidy, k_idx, bidz]
# tAgSFA = gSFA_mkl[tidx, None, bidx, k_idx, bidz]
# tBgSFB = gSFB_nkl[0, None, bidy, k_idx, bidz]
# # Prefetch for next iteration to hide memory latency
# if unroll_idx == 0 and k_tile + unroll_factor < k_tile_cnt:
# tAgA_next = gA_mkl[tidx, None, bidx, k_tile + unroll_factor, bidz]
# tBgB_next = gB_nkl[0, None, bidy, k_tile + unroll_factor, bidz]
# _ = tAgA_next.load() # Prefetch
# _ = tBgB_next.load()
# tArA = cute.make_rmem_tensor_like(tAgA, cutlass.Float32)
# tBrB = cute.make_rmem_tensor_like(tBgB, cutlass.Float32)
# tArSFA = cute.make_rmem_tensor_like(tAgSFA, cutlass.Float32)
# tBrSFB = cute.make_rmem_tensor_like(tBgSFB, cutlass.Float32)
# # Asynchronous loads with conversion pipeline
# a_val_nvfp4 = tAgA.load()
# b_val_nvfp4 = tBgB.load()
# sfa_val_fp8 = tAgSFA.load()
# sfb_val_fp8 = tBgSFB.load()
# # Convert and store
# tArA.store(a_val_nvfp4.to(cutlass.Float32))
# tBrB.store(b_val_nvfp4.to(cutlass.Float32))
# tArSFA.store(sfa_val_fp8.to(cutlass.Float32))
# tBrSFB.store(sfb_val_fp8.to(cutlass.Float32))
# # Fused multiply-accumulate
# for i in cutlass.range_constexpr(mma_tiler_mnk[2]):
# res += tArA[i] * tArSFA[i] * tBrB[i] * tBrSFB[i]
# # Handle remaining k-tiles
# for k_tile in range(k_tile_unrolled, k_tile_cnt):
# tAgA = gA_mkl[tidx, None, bidx, k_tile, bidz]
# tBgB = gB_nkl[0, None, bidy, k_tile, bidz]
# tAgSFA = gSFA_mkl[tidx, None, bidx, k_tile, bidz]
# tBgSFB = gSFB_nkl[0, None, bidy, k_tile, bidz]
# tArA = cute.make_rmem_tensor_like(tAgA, cutlass.Float32)
# tBrB = cute.make_rmem_tensor_like(tBgB, cutlass.Float32)
# tArSFA = cute.make_rmem_tensor_like(tAgSFA, cutlass.Float32)
# tBrSFB = cute.make_rmem_tensor_like(tBgSFB, cutlass.Float32)
# a_val_nvfp4 = tAgA.load()
# b_val_nvfp4 = tBgB.load()
# sfa_val_fp8 = tAgSFA.load()
# sfb_val_fp8 = tBgSFB.load()
# tArA.store(a_val_nvfp4.to(cutlass.Float32))
# tBrB.store(b_val_nvfp4.to(cutlass.Float32))
# tArSFA.store(sfa_val_fp8.to(cutlass.Float32))
# tBrSFB.store(sfb_val_fp8.to(cutlass.Float32))
# for i in cutlass.range_constexpr(mma_tiler_mnk[2]):
# res += tArA[i] * tArSFA[i] * tBrB[i] * tBrSFB[i]
# # Store with async copy for better throughput
# tCgC.store(res.to(cutlass.Float16))
# return
# @cute.jit
# def my_kernel(
# 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)
# # Optimized grid for Blackwell: larger tiles = fewer blocks
# grid = (
# cute.ceil_div(c_tensor.shape[0], 256), # Match larger M tile size
# 1,
# c_tensor.shape[2],
# )
# # Launch with increased thread count and persistent kernel hint
# kernel(a_tensor, b_tensor, sfa_tensor, sfb_tensor, c_tensor).launch(
# grid=grid,
# block=[threads_per_cta, 1, 1],
# cluster=(1, 1, 1),
# shared_memory=0, # Explicitly set for Blackwell's larger shared memory
# )
# return
# _compiled_kernel_cache = None
# def compile_kernel():
# global _compiled_kernel_cache
# if _compiled_kernel_cache is not None:
# return _compiled_kernel_cache
# 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)
# 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)
# # Compile with Blackwell-specific optimizations
# _compiled_kernel_cache = cute.compile(
# my_kernel,
# a_ptr, b_ptr, sfa_ptr, sfb_ptr, c_ptr, (0, 0, 0, 0),
# # Add Blackwell-specific compiler flags if available
# # compile_options={'arch': 'sm_100', 'maxrregcount': 255}
# )
# return _compiled_kernel_cache
# def custom_kernel(data: input_t) -> output_t:
# a, b, _, _, sfa_permuted, sfb_permuted, c = data
# compiled_func = compile_kernel()
# m, k, l = a.shape
# k = k * 2
# n = 1
# 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
# )
# compiled_func(a_ptr, b_ptr, sfa_ptr, sfb_ptr, c_ptr, (m, n, k, l))
# return c
# 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.blockscaled_layout as blockscaled_utils
# # Blackwell-optimized kernel configuration
# mma_tiler_mnk = (256, 1, 128) # Larger tiles for Blackwell's increased SM count and register file
# ab_dtype = cutlass.Float4E2M1FN
# sf_dtype = cutlass.Float8E4M3FN
# c_dtype = cutlass.Float16
# sf_vec_size = 16
# threads_per_cta = 256 # Increased thread count for better occupancy on Blackwell
# def ceil_div(a, b):
# return (a + b - 1) // b
# @cute.kernel
# def kernel(
# mA_mkl: cute.Tensor,
# mB_nkl: cute.Tensor,
# mSFA_mkl: cute.Tensor,
# mSFB_nkl: cute.Tensor,
# mC_mnl: cute.Tensor,
# ):
# bidx, bidy, bidz = cute.arch.block_idx()
# tidx, _, _ = cute.arch.thread_idx()
# # Local tile extraction with larger tile sizes
# 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)
# )
# tCgC = gC_mnl[tidx, None, bidx, bidy, bidz]
# tCgC = cute.make_tensor(tCgC.iterator, 1)
# # Use FP32 accumulation for better precision on Blackwell
# res = cute.zeros_like(tCgC, cutlass.Float32)
# k_tile_cnt = gA_mkl.layout[3].shape
# # Unroll factor optimized for Blackwell's instruction cache and pipeline depth
# unroll_factor = 8 # Increased from 4 for better ILP
# k_tile_unrolled = (k_tile_cnt // unroll_factor) * unroll_factor
# # Main unrolled loop for better ILP on Blackwell
# for k_tile in cutlass.range_constexpr(0, k_tile_unrolled, unroll_factor):
# # Prefetch multiple iterations ahead
# if k_tile + unroll_factor < k_tile_cnt:
# for prefetch_offset in cutlass.range_constexpr(unroll_factor):
# tAgA_pf = gA_mkl[tidx, None, bidx, k_tile + unroll_factor + prefetch_offset, bidz]
# tBgB_pf = gB_nkl[0, None, bidy, k_tile + unroll_factor + prefetch_offset, bidz]
# _ = tAgA_pf.load()
# _ = tBgB_pf.load()
# # Process unrolled k-tiles
# for unroll_idx in cutlass.range_constexpr(unroll_factor):
# k_idx = k_tile + unroll_idx
# tAgA = gA_mkl[tidx, None, bidx, k_idx, bidz]
# tBgB = gB_nkl[0, None, bidy, k_idx, bidz]
# tAgSFA = gSFA_mkl[tidx, None, bidx, k_idx, bidz]
# tBgSFB = gSFB_nkl[0, None, bidy, k_idx, bidz]
# tArA = cute.make_rmem_tensor_like(tAgA, cutlass.Float32)
# tBrB = cute.make_rmem_tensor_like(tBgB, cutlass.Float32)
# tArSFA = cute.make_rmem_tensor_like(tAgSFA, cutlass.Float32)
# tBrSFB = cute.make_rmem_tensor_like(tBgSFB, cutlass.Float32)
# # Load and convert in pipeline
# a_val_nvfp4 = tAgA.load()
# b_val_nvfp4 = tBgB.load()
# sfa_val_fp8 = tAgSFA.load()
# sfb_val_fp8 = tBgSFB.load()
# tArA.store(a_val_nvfp4.to(cutlass.Float32))
# tBrB.store(b_val_nvfp4.to(cutlass.Float32))
# tArSFA.store(sfa_val_fp8.to(cutlass.Float32))
# tBrSFB.store(sfb_val_fp8.to(cutlass.Float32))
# # Fused multiply-accumulate with 2x unroll
# for i in cutlass.range_constexpr(0, mma_tiler_mnk[2], 2):
# res += tArA[i] * tArSFA[i] * tBrB[i] * tBrSFB[i]
# if i + 1 < mma_tiler_mnk[2]:
# res += tArA[i+1] * tArSFA[i+1] * tBrB[i+1] * tBrSFB[i+1]
# # Handle remaining k-tiles
# for k_tile in range(k_tile_unrolled, k_tile_cnt):
# tAgA = gA_mkl[tidx, None, bidx, k_tile, bidz]
# tBgB = gB_nkl[0, None, bidy, k_tile, bidz]
# tAgSFA = gSFA_mkl[tidx, None, bidx, k_tile, bidz]
# tBgSFB = gSFB_nkl[0, None, bidy, k_tile, bidz]
# tArA = cute.make_rmem_tensor_like(tAgA, cutlass.Float32)
# tBrB = cute.make_rmem_tensor_like(tBgB, cutlass.Float32)
# tArSFA = cute.make_rmem_tensor_like(tAgSFA, cutlass.Float32)
# tBrSFB = cute.make_rmem_tensor_like(tBgSFB, cutlass.Float32)
# a_val_nvfp4 = tAgA.load()
# b_val_nvfp4 = tBgB.load()
# sfa_val_fp8 = tAgSFA.load()
# sfb_val_fp8 = tBgSFB.load()
# tArA.store(a_val_nvfp4.to(cutlass.Float32))
# tBrB.store(b_val_nvfp4.to(cutlass.Float32))
# tArSFA.store(sfa_val_fp8.to(cutlass.Float32))
# tBrSFB.store(sfb_val_fp8.to(cutlass.Float32))
# for i in cutlass.range_constexpr(mma_tiler_mnk[2]):
# res += tArA[i] * tArSFA[i] * tBrB[i] * tBrSFB[i]
# # Store with conversion
# tCgC.store(res.to(cutlass.Float16))
# return
# @cute.jit
# def my_kernel(
# 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)
# # Optimized grid for Blackwell: larger tiles = fewer blocks
# grid = (
# cute.ceil_div(c_tensor.shape[0], 256), # Match larger M tile size
# 1,
# c_tensor.shape[2],
# )
# # Launch with increased thread count - removed invalid 'shared_memory' kwarg
# kernel(a_tensor, b_tensor, sfa_tensor, sfb_tensor, c_tensor).launch(
# grid=grid,
# block=[threads_per_cta, 1, 1],
# cluster=(1, 1, 1),
# )
# return
# _compiled_kernel_cache = None
# def compile_kernel():
# global _compiled_kernel_cache
# if _compiled_kernel_cache is not None:
# return _compiled_kernel_cache
# 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)
# 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)
# # Compile kernel - removed invalid compile_options
# _compiled_kernel_cache = cute.compile(
# my_kernel,
# a_ptr, b_ptr, sfa_ptr, sfb_ptr, c_ptr, (0, 0, 0, 0),
# )
# return _compiled_kernel_cache
# def custom_kernel(data: input_t) -> output_t:
# a, b, _, _, sfa_permuted, sfb_permuted, c = data
# compiled_func = compile_kernel()
# m, k, l = a.shape
# k = k * 2
# n = 1
# 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
# )
# compiled_func(a_ptr, b_ptr, sfa_ptr, sfb_ptr, c_ptr, (m, n, k, l))
# return c
# 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.blockscaled_layout as blockscaled_utils
# # Blackwell-optimized kernel configuration
# mma_tiler_mnk = (256, 1, 128) # Larger tiles for Blackwell's increased SM count and register file
# ab_dtype = cutlass.Float4E2M1FN
# sf_dtype = cutlass.Float8E4M3FN
# c_dtype = cutlass.Float16
# sf_vec_size = 16
# threads_per_cta = 256 # Increased thread count for better occupancy on Blackwell
# def ceil_div(a, b):
# return (a + b - 1) // b
# @cute.kernel
# def kernel(
# mA_mkl: cute.Tensor,
# mB_nkl: cute.Tensor,
# mSFA_mkl: cute.Tensor,
# mSFB_nkl: cute.Tensor,
# mC_mnl: cute.Tensor,
# ):
# bidx, bidy, bidz = cute.arch.block_idx()
# tidx, _, _ = cute.arch.thread_idx()
# # Local tile extraction with larger tile sizes
# 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)
# )
# tCgC = gC_mnl[tidx, None, bidx, bidy, bidz]
# tCgC = cute.make_tensor(tCgC.iterator, 1)
# # Use FP32 accumulation for better precision on Blackwell
# res = cute.zeros_like(tCgC, cutlass.Float32)
# k_tile_cnt = gA_mkl.layout[3].shape
# # Process all k-tiles with manual unrolling where beneficial
# for k_tile in range(k_tile_cnt):
# tAgA = gA_mkl[tidx, None, bidx, k_tile, bidz]
# tBgB = gB_nkl[0, None, bidy, k_tile, bidz]
# tAgSFA = gSFA_mkl[tidx, None, bidx, k_tile, bidz]
# tBgSFB = gSFB_nkl[0, None, bidy, k_tile, bidz]
# tArA = cute.make_rmem_tensor_like(tAgA, cutlass.Float32)
# tBrB = cute.make_rmem_tensor_like(tBgB, cutlass.Float32)
# tArSFA = cute.make_rmem_tensor_like(tAgSFA, cutlass.Float32)
# tBrSFB = cute.make_rmem_tensor_like(tBgSFB, cutlass.Float32)
# # Load and convert in pipeline
# a_val_nvfp4 = tAgA.load()
# b_val_nvfp4 = tBgB.load()
# sfa_val_fp8 = tAgSFA.load()
# sfb_val_fp8 = tBgSFB.load()
# tArA.store(a_val_nvfp4.to(cutlass.Float32))
# tBrB.store(b_val_nvfp4.to(cutlass.Float32))
# tArSFA.store(sfa_val_fp8.to(cutlass.Float32))
# tBrSFB.store(sfb_val_fp8.to(cutlass.Float32))
# # Fused multiply-accumulate - use constexpr for compile-time constant
# for i in cutlass.range_constexpr(mma_tiler_mnk[2]):
# res += tArA[i] * tArSFA[i] * tBrB[i] * tBrSFB[i]
# # Store with conversion
# tCgC.store(res.to(cutlass.Float16))
# return
# @cute.jit
# def my_kernel(
# 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)
# # Optimized grid for Blackwell: larger tiles = fewer blocks
# grid = (
# cute.ceil_div(c_tensor.shape[0], mma_tiler_mnk[0]),
# 1,
# c_tensor.shape[2],
# )
# # Launch with increased thread count
# kernel(a_tensor, b_tensor, sfa_tensor, sfb_tensor, c_tensor).launch(
# grid=grid,
# block=[threads_per_cta, 1, 1],
# cluster=(1, 1, 1),
# )
# return
# _compiled_kernel_cache = None
# def compile_kernel():
# global _compiled_kernel_cache
# if _compiled_kernel_cache is not None:
# return _compiled_kernel_cache
# 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)
# 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)
# # Compile kernel
# _compiled_kernel_cache = cute.compile(
# my_kernel,
# a_ptr, b_ptr, sfa_ptr, sfb_ptr, c_ptr, (0, 0, 0, 0),
# )
# return _compiled_kernel_cache
# def custom_kernel(data: input_t) -> output_t:
# a, b, _, _, sfa_permuted, sfb_permuted, c = data
# compiled_func = compile_kernel()
# m, k, l = a.shape
# k = k * 2
# n = 1
# 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
# )
# compiled_func(a_ptr, b_ptr, sfa_ptr, sfb_ptr, c_ptr, (m, n, k, l))
# return c
# 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.blockscaled_layout as blockscaled_utils
# # Blackwell-optimized kernel configuration
# mma_tiler_mnk = (128, 1, 64) # Conservative tile sizes for correctness first
# ab_dtype = cutlass.Float4E2M1FN
# sf_dtype = cutlass.Float8E4M3FN
# c_dtype = cutlass.Float16
# sf_vec_size = 16
# threads_per_cta = 128 # Match original for stability
# def ceil_div(a, b):
# return (a + b - 1) // b
# @cute.kernel
# def kernel(
# mA_mkl: cute.Tensor,
# mB_nkl: cute.Tensor,
# mSFA_mkl: cute.Tensor,
# mSFB_nkl: cute.Tensor,
# mC_mnl: cute.Tensor,
# ):
# bidx, bidy, bidz = cute.arch.block_idx()
# tidx, _, _ = cute.arch.thread_idx()
# # Local tile extraction
# gA_mkl = cute.local_tile(
# mA_mkl, cute.slice_(mma_tiler_mnk, (None, 0, None)), (bidx, None, bidz)
# )
# gSFA_mkl = cute.local_tile(
# mSFA_mkl, cute.slice_(mma_tiler_mnk, (None, 0, None)), (bidx, None, bidz)
# )
# gB_nkl = cute.local_tile(
# mB_nkl, cute.slice_(mma_tiler_mnk, (0, None, None)), (bidy, None, bidz)
# )
# gSFB_nkl = cute.local_tile(
# mSFB_nkl, cute.slice_(mma_tiler_mnk, (0, None, None)), (bidy, None, bidz)
# )
# gC_mnl = cute.local_tile(
# mC_mnl, cute.slice_(mma_tiler_mnk, (None, None, 0)), (bidx, bidy, bidz)
# )
# tCgC = gC_mnl[tidx, None]
# tCgC = cute.make_tensor(tCgC.iterator, 1)
# # Use FP32 accumulation
# res = cute.zeros_like(tCgC, cutlass.Float32)
# k_tile_cnt = gA_mkl.layout[1].shape
# # Process all k-tiles
# for k_tile in range(k_tile_cnt):
# tAgA = gA_mkl[tidx, None, k_tile]
# tBgB = gB_nkl[0, None, k_tile]
# tAgSFA = gSFA_mkl[tidx, None, k_tile]
# tBgSFB = gSFB_nkl[0, None, k_tile]
# tArA = cute.make_rmem_tensor_like(tAgA, cutlass.Float32)
# tBrB = cute.make_rmem_tensor_like(tBgB, cutlass.Float32)
# tArSFA = cute.make_rmem_tensor_like(tAgSFA, cutlass.Float32)
# tBrSFB = cute.make_rmem_tensor_like(tBgSFB, cutlass.Float32)
# # Load and convert
# a_val_nvfp4 = tAgA.load()
# b_val_nvfp4 = tBgB.load()
# sfa_val_fp8 = tAgSFA.load()
# sfb_val_fp8 = tBgSFB.load()
# tArA.store(a_val_nvfp4.to(cutlass.Float32))
# tBrB.store(b_val_nvfp4.to(cutlass.Float32))
# tArSFA.store(sfa_val_fp8.to(cutlass.Float32))
# tBrSFB.store(sfb_val_fp8.to(cutlass.Float32))
# # Fused multiply-accumulate
# for i in cutlass.range_constexpr(mma_tiler_mnk[2]):
# res += tArA[i] * tArSFA[i] * tBrB[i] * tBrSFB[i]
# # Store result
# tCgC.store(res.to(cutlass.Float16))
# return
# @cute.jit
# def my_kernel(
# 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 calculation
# grid = (
# cute.ceil_div(c_tensor.shape[0], mma_tiler_mnk[0]),
# 1,
# c_tensor.shape[2],
# )
# # Launch kernel
# kernel(a_tensor, b_tensor, sfa_tensor, sfb_tensor, c_tensor).launch(
# grid=grid,
# block=[threads_per_cta, 1, 1],
# cluster=(1, 1, 1),
# )
# return
# _compiled_kernel_cache = None
# def compile_kernel():
# global _compiled_kernel_cache
# if _compiled_kernel_cache is not None:
# return _compiled_kernel_cache
# 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)
# 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_kernel_cache = cute.compile(
# my_kernel,
# a_ptr, b_ptr, sfa_ptr, sfb_ptr, c_ptr, (0, 0, 0, 0),
# )
# return _compiled_kernel_cache
# def custom_kernel(data: input_t) -> output_t:
# a, b, _, _, sfa_permuted, sfb_permuted, c = data
# compiled_func = compile_kernel()
# m, k, l = a.shape
# k = k * 2
# n = 1
# 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
# )
# compiled_func(a_ptr, b_ptr, sfa_ptr, sfb_ptr, c_ptr, (m, n, k, l))
# return c
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.blockscaled_layout as blockscaled_utils
# Use reference kernel configuration for correctness
mma_tiler_mnk = (128, 1, 64)
ab_dtype = cutlass.Float4E2M1FN
sf_dtype = cutlass.Float8E4M3FN
c_dtype = cutlass.Float16
sf_vec_size = 16
threads_per_cta = 128
def ceil_div(a, b):
return (a + b - 1) // b
@cute.kernel
def kernel(
mA_mkl: cute.Tensor,
mB_nkl: cute.Tensor,
mSFA_mkl: cute.Tensor,
mSFB_nkl: cute.Tensor,
mC_mnl: cute.Tensor,
):
bidx, bidy, bidz = cute.arch.block_idx()
tidx, _, _ = cute.arch.thread_idx()
# Extract global tiles - match reference implementation exactly
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)
)
# Index into output - use reference pattern
tCgC = gC_mnl[tidx, None, bidx, bidy, bidz]
tCgC = cute.make_tensor(tCgC.iterator, 1)
# Accumulate in FP32
res = cute.zeros_like(tCgC, cutlass.Float32)
# Get k-tile count
k_tile_cnt = gA_mkl.layout[3].shape
# Main computation loop
for k_tile in range(k_tile_cnt):
# Index A, B and scale factors for this k-tile
tAgA = gA_mkl[tidx, None, bidx, k_tile, bidz]
tBgB = gB_nkl[0, None, bidy, k_tile, bidz]
tAgSFA = gSFA_mkl[tidx, None, bidx, k_tile, bidz]
tBgSFB = gSFB_nkl[0, None, bidy, k_tile, bidz]
# Create register tensors
tArA = cute.make_rmem_tensor_like(tAgA, cutlass.Float32)
tBrB = cute.make_rmem_tensor_like(tBgB, cutlass.Float32)
tArSFA = cute.make_rmem_tensor_like(tAgSFA, cutlass.Float32)
tBrSFB = cute.make_rmem_tensor_like(tBgSFB, cutlass.Float32)
# Load from global memory
a_val_nvfp4 = tAgA.load()
b_val_nvfp4 = tBgB.load()
sfa_val_fp8 = tAgSFA.load()
sfb_val_fp8 = tBgSFB.load()
# Convert and store to registers
tArA.store(a_val_nvfp4.to(cutlass.Float32))
tBrB.store(b_val_nvfp4.to(cutlass.Float32))
tArSFA.store(sfa_val_fp8.to(cutlass.Float32))
tBrSFB.store(sfb_val_fp8.to(cutlass.Float32))
# Compute: accumulate scaled dot product
for i in cutlass.range_constexpr(mma_tiler_mnk[2]):
res += tArA[i] * tArSFA[i] * tBrB[i] * tBrSFB[i]
# Convert back to FP16 and store
tCgC.store(res.to(cutlass.Float16))
return
@cute.jit
def my_kernel(
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
# Create tensor views with proper layouts
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))
)
# Create scale factor tensor layouts
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)
# Launch grid
grid = (
cute.ceil_div(c_tensor.shape[0], mma_tiler_mnk[0]),
1,
c_tensor.shape[2],
)
# Launch kernel
kernel(a_tensor, b_tensor, sfa_tensor, sfb_tensor, c_tensor).launch(
grid=grid,
block=[threads_per_cta, 1, 1],
cluster=(1, 1, 1),
)
return
_compiled_kernel_cache = None
def compile_kernel():
global _compiled_kernel_cache
if _compiled_kernel_cache is not None:
return _compiled_kernel_cache
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)
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_kernel_cache = cute.compile(
my_kernel,
a_ptr, b_ptr, sfa_ptr, sfb_ptr, c_ptr, (0, 0, 0, 0),
)
return _compiled_kernel_cache
def custom_kernel(data: input_t) -> output_t:
a, b, _, _, sfa_permuted, sfb_permuted, c = data
compiled_func = compile_kernel()
m, k, l = a.shape
k = k * 2
n = 1
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
)
compiled_func(a_ptr, b_ptr, sfa_ptr, sfb_ptr, c_ptr, (m, n, k, l))
return cscrolls · 1053 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