submission 116315
swanbomb_ · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 730 lines, June 9 Researcher Reciprocity License v1.0.
lversion.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-nvfp4-gemv-116315?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:1a32cb7d640e14f1e61dce24e3ec164068f00bb4d249766368bf49dae8f7ccb3
license declaredunknown
license concludedunknown
authorsswanbomb_
imported2026-08-15
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
shared-memory
def warp_smem_reduce(threads_m: int, threads_k: int, threads_l):Kernel source
lversion.py730 lines
import cutlass
from enum import Enum
from cutlass import Float32, Float16, Int16, Int32, Int8
import cutlass.cute as cute
import cutlass.utils.blockscaled_layout as blockscaled_utils
from cutlass.cute import arch
from cutlass.cute.arch import builtin, arith, ir, llvm, vector, dsl_user_op, T
from cutlass.cute.tensor import TensorSSA
from cutlass.cute.core import slice_
from cutlass.cute.runtime import make_ptr
from cutlass.cute.typing import Pointer
from task import input_t, output_t
ab_dtype = cutlass.Float4E2M1FN
sf_dtype = cutlass.Float8E4M3FN
sf_vec_size = 16
acc_dtype = Float32
c_dtype = Float16
sfc_dtype = Float16
def ceil_div(a, b):
return (a + b - 1) // b
# from FlashAttention repo
@cute.jit
def scalar_to_ssa(a: cute.Numeric, dtype) -> cute.TensorSSA:
vec = cute.make_rmem_tensor(1, dtype)
vec[0] = a
return vec.load()
# https://veitner.bearblog.dev/demystifying-numeric-conversions-in-cutedsl/
@dsl_user_op
def cvt_f8e4m3_f16_intr(vec_f8e4m3, length, *, loc=None, ip=None):
src_pos = 0
vec_src_i8 = builtin.unrealized_conversion_cast(
[ir.VectorType.get([length], Int8.mlir_type, loc=loc)],
[vec_f8e4m3],
loc=loc,
ip=ip,
)
vec_i8x8_type = ir.VectorType.get([8], Int8.mlir_type, loc=loc)
vec_i8x4_type = ir.VectorType.get([4], Int8.mlir_type, loc=loc)
vec_i8x2_type = ir.VectorType.get([2], Int8.mlir_type, loc=loc)
vec_dst_type = ir.VectorType.get([length], Float16.mlir_type, loc=loc)
vec_dst = llvm.mlir_zero(vec_dst_type, loc=loc, ip=ip)
# try to use vectorized version
if length >= 8:
num_vec8 = length // 8
for _ in range(num_vec8):
vec_f8e4m3x8 = vector.extract_strided_slice(
vec_i8x8_type, vec_src_i8, [src_pos], [8], [1], loc=loc, ip=ip
)
vec_f16x8 = cvt_f8e4m3x8_to_f16x8(vec_f8e4m3x8, loc=loc, ip=ip)
vec_dst = vector.insert_strided_slice(
vec_f16x8, vec_dst, [src_pos], [1], loc=loc, ip=ip
)
src_pos += 8
length -= 8
if length >= 4:
vec_f8e4m3x4 = vector.extract_strided_slice(
vec_i8x4_type, vec_src_i8, [src_pos], [4], [1], loc=loc, ip=ip
)
vec_f16x4 = cvt_f8e4m3x4_to_f16x4(vec_f8e4m3x4, loc=loc, ip=ip)
vec_dst = vector.insert_strided_slice(vec_f16x4, vec_dst, [src_pos], [1], loc=loc, ip=ip)
src_pos += 4
length -= 4
if length >= 2:
vec_f8e4m3x2 = vector.extract_strided_slice(
vec_i8x2_type, vec_src_i8, [src_pos], [2], [1], loc=loc, ip=ip
)
vec_f16x2 = cvt_f8e4m3x2_to_f16x2(vec_f8e4m3x2, loc=loc, ip=ip)
vec_dst = vector.insert_strided_slice(vec_f16x2, vec_dst, [src_pos], [1], loc=loc, ip=ip)
src_pos += 2
length -= 2
if length >= 1:
val_f16 = cvt_f8e4m3_f16(
vector.extractelement(
vec_src_i8,
position=arith.constant(Int32.mlir_type, src_pos),
loc=loc,
ip=ip,
),
loc=loc,
ip=ip,
)
vec_dst = vector.insertelement(
val_f16,
vec_dst,
position=arith.constant(Int32.mlir_type, src_pos),
loc=loc,
ip=ip,
)
return vec_dst
@dsl_user_op
def cvt_f8e4m3_f16(src, *, loc=None, ip=None):
# 0 padding for upper 8 bits
zero = arith.constant(src.type, 0, loc=loc, ip=ip)
vec2 = vector.from_elements(
ir.VectorType.get([2], src.type, loc=loc), [src, zero], loc=loc, ip=ip
)
rst_vec2 = cvt_f8e4m3x2_to_f16x2(vec2, loc=loc, ip=ip)
# only the 1st element is valid
rst = vector.extract(rst_vec2, dynamic_position=[], static_position=[0], loc=loc, ip=ip)
return rst
# Convert 2 float8e4m3 values to 2 float16 values
@dsl_user_op
def cvt_f8e4m3x2_to_f16x2(src_vec2, *, loc=None, ip=None):
# pack 2 float8e4m3 into 1 int16 value
src_i16 = llvm.bitcast(Int16.mlir_type, src_vec2, loc=loc, ip=ip)
rst_i32 = llvm.inline_asm(
Int32.mlir_type,
[src_i16],
"""{\n\t
cvt.rn.f16x2.e4m3x2 $0, $1;\n\t
}""",
"=r,h",
)
vec_f16x2_type = ir.VectorType.get([2], Float16.mlir_type, loc=loc)
vec_f16x2 = llvm.bitcast(vec_f16x2_type, rst_i32, loc=loc, ip=ip)
return vec_f16x2
# Convert 4 float8e4m3 values to 4 float16 values
@dsl_user_op
def cvt_f8e4m3x4_to_f16x4(src_vec4, *, loc=None, ip=None):
# pack 4 float8e4m3 into 1 int32 value
src_i32 = llvm.bitcast(Int32.mlir_type, src_vec4, loc=loc, ip=ip)
rst_i32x2 = llvm.inline_asm(
llvm.StructType.get_literal([T.i32(), T.i32()]),
[src_i32],
"""{\n\t
.reg .b16 h0, h1;\n\t
mov.b32 {h0, h1}, $2;\n\t
cvt.rn.f16x2.e4m3x2 $0, h0;\n\t
cvt.rn.f16x2.e4m3x2 $1, h1;\n\t
}""",
"=r,=r,r",
)
res0 = llvm.extractvalue(T.i32(), rst_i32x2, [0])
res1 = llvm.extractvalue(T.i32(), rst_i32x2, [1])
vec_i32x2_type = ir.VectorType.get([2], Int32.mlir_type, loc=loc)
vec_i32x2 = vector.from_elements(vec_i32x2_type, [res0, res1], loc=loc, ip=ip)
vec_f16x4_type = ir.VectorType.get([4], Float16.mlir_type, loc=loc)
vec_f16x4 = llvm.bitcast(vec_f16x4_type, vec_i32x2, loc=loc, ip=ip)
return vec_f16x4
# Convert 8 float8e4m3 values to 8 float16 values
@dsl_user_op
def cvt_f8e4m3x8_to_f16x8(src_vec8, *, loc=None, ip=None):
# Split into two i32 values instead of using i64
vec_i32x2_type = ir.VectorType.get([2], Int32.mlir_type, loc=loc)
src_i32x2 = llvm.bitcast(vec_i32x2_type, src_vec8, loc=loc, ip=ip)
src_lo = llvm.extractelement(src_i32x2, arith.constant(Int32.mlir_type, 0), loc=loc, ip=ip)
src_hi = llvm.extractelement(src_i32x2, arith.constant(Int32.mlir_type, 1), loc=loc, ip=ip)
# Process lower 4 bytes (4 fp8 values)
rst_lo_i32x2 = llvm.inline_asm(
llvm.StructType.get_literal([T.i32(), T.i32()]),
[src_lo],
"""{\n\t
.reg .b16 h0, h1;\n\t
mov.b32 {h0, h1}, $2;\n\t
cvt.rn.f16x2.e4m3x2 $0, h0;\n\t
cvt.rn.f16x2.e4m3x2 $1, h1;\n\t
}""",
"=r,=r,r",
)
# Process upper 4 bytes (4 fp8 values)
rst_hi_i32x2 = llvm.inline_asm(
llvm.StructType.get_literal([T.i32(), T.i32()]),
[src_hi],
"""{\n\t
.reg .b16 h0, h1;\n\t
mov.b32 {h0, h1}, $2;\n\t
cvt.rn.f16x2.e4m3x2 $0, h0;\n\t
cvt.rn.f16x2.e4m3x2 $1, h1;\n\t
}""",
"=r,=r,r",
)
res0 = llvm.extractvalue(T.i32(), rst_lo_i32x2, [0])
res1 = llvm.extractvalue(T.i32(), rst_lo_i32x2, [1])
res2 = llvm.extractvalue(T.i32(), rst_hi_i32x2, [0])
res3 = llvm.extractvalue(T.i32(), rst_hi_i32x2, [1])
vec_i32x4_type = ir.VectorType.get([4], Int32.mlir_type, loc=loc)
vec_i32x4 = vector.from_elements(vec_i32x4_type, [res0, res1, res2, res3], loc=loc, ip=ip)
vec_f16x8_type = ir.VectorType.get([8], Float16.mlir_type, loc=loc)
vec_f16x8 = llvm.bitcast(vec_f16x8_type, vec_i32x4, loc=loc, ip=ip)
return vec_f16x8
@dsl_user_op
def fma_f16x2(
a: tuple[Float16, Float16],
b: tuple[Float16, Float16],
c: tuple[Float16, Float16],
*,
loc=None,
ip=None,
) -> tuple[Float16, Float16]:
# Pack two Float16 values into vector<2xf16>
vec_type = ir.VectorType.get([2], Float16.mlir_type, loc=loc)
vec_a = vector.from_elements(
vec_type,
[a[0].ir_value(loc=loc, ip=ip), a[1].ir_value(loc=loc, ip=ip)],
loc=loc,
ip=ip,
)
vec_b = vector.from_elements(
vec_type,
[b[0].ir_value(loc=loc, ip=ip), b[1].ir_value(loc=loc, ip=ip)],
loc=loc,
ip=ip,
)
vec_c = vector.from_elements(
vec_type,
[c[0].ir_value(loc=loc, ip=ip), c[1].ir_value(loc=loc, ip=ip)],
loc=loc,
ip=ip,
)
# Bitcast to i32 for PTX (f16x2 is packed into 32 bits)
a_i32 = llvm.bitcast(Int32.mlir_type, vec_a, loc=loc, ip=ip)
b_i32 = llvm.bitcast(Int32.mlir_type, vec_b, loc=loc, ip=ip)
c_i32 = llvm.bitcast(Int32.mlir_type, vec_c, loc=loc, ip=ip)
result_i32 = llvm.inline_asm(
Int32.mlir_type,
[a_i32, b_i32, c_i32],
"fma.rn.f16x2 $0, $1, $2, $3;",
"=r,r,r,r",
has_side_effects=False,
is_align_stack=False,
asm_dialect=llvm.AsmDialect.AD_ATT,
loc=loc,
ip=ip,
)
# Bitcast back to vector<2xf16>
vec_result = llvm.bitcast(vec_type, result_i32, loc=loc, ip=ip)
# Extract results
result0 = Float16(vector.extract(vec_result, [], [0], loc=loc, ip=ip))
result1 = Float16(vector.extract(vec_result, [], [1], loc=loc, ip=ip))
return result0, result1
@dsl_user_op
def binary_f16x2(
a: tuple[Float16, Float16], b: tuple[Float16, Float16], asm_string: str, *, loc=None, ip=None
) -> tuple[Float16, Float16]:
# Pack two Float16 values into vector<2xf16>
vec_type = ir.VectorType.get([2], Float16.mlir_type, loc=loc)
vec_a = vector.from_elements(
vec_type,
[a[0].ir_value(loc=loc, ip=ip), a[1].ir_value(loc=loc, ip=ip)],
loc=loc,
ip=ip,
)
vec_b = vector.from_elements(
vec_type,
[b[0].ir_value(loc=loc, ip=ip), b[1].ir_value(loc=loc, ip=ip)],
loc=loc,
ip=ip,
)
# Bitcast to i32 for PTX (f16x2 is packed into 32 bits)
a_i32 = llvm.bitcast(Int32.mlir_type, vec_a, loc=loc, ip=ip)
b_i32 = llvm.bitcast(Int32.mlir_type, vec_b, loc=loc, ip=ip)
result_i32 = llvm.inline_asm(
Int32.mlir_type,
[a_i32, b_i32],
asm_string,
"=r,r,r",
has_side_effects=False,
is_align_stack=False,
asm_dialect=llvm.AsmDialect.AD_ATT,
loc=loc,
ip=ip,
)
# Bitcast back to vector<2xf16>
vec_result = llvm.bitcast(vec_type, result_i32, loc=loc, ip=ip)
# Extract results
result0 = Float16(vector.extract(vec_result, [], [0], loc=loc, ip=ip))
result1 = Float16(vector.extract(vec_result, [], [1], loc=loc, ip=ip))
return result0, result1
@dsl_user_op
def add_f16x2(
a: tuple[Float16, Float16], b: tuple[Float16, Float16], *, loc=None, ip=None
) -> tuple[Float16, Float16]:
return binary_f16x2(a, b, "add.f16x2 $0, $1, $2;", loc=loc, ip=ip)
@dsl_user_op
def mul_f16x2(
a: tuple[Float16, Float16], b: tuple[Float16, Float16], *, loc=None, ip=None
) -> tuple[Float16, Float16]:
return binary_f16x2(a, b, "mul.f16x2 $0, $1, $2;", loc=loc, ip=ip)
@cute.jit
def make_tensors(
a_ptr: Pointer,
b_ptr: Pointer,
sfa_ptr: Pointer,
sfb_ptr: Pointer,
c_ptr: Pointer,
size: tuple[int, int, int, int],
) -> tuple[cute.Tensor, cute.Tensor, cute.Tensor, cute.Tensor, cute.Tensor]:
m, _, k, l = size
a_tensor = cute.make_tensor(
a_ptr,
cute.make_layout(
(cute.assume(m, 64), cute.assume(k, 64), cute.assume(l, 64)),
stride=(cute.assume(k, 64), 1, cute.assume(m * k, 64)),
),
)
n_padded = 128
b_tensor = cute.make_tensor(
b_ptr,
cute.make_layout(
(n_padded, cute.assume(k, 64), cute.assume(l, 64)),
stride=(cute.assume(k, 64), 1, cute.assume(n_padded * k, 64)),
),
)
c_tensor = cute.make_tensor(
c_ptr,
cute.make_layout(
(cute.assume(m, 64), 1, cute.assume(l, 64)),
stride=(1, 1, cute.assume(m, 64)),
),
)
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)
return (a_tensor, b_tensor, sfa_tensor, sfb_tensor, c_tensor)
def warp_smem_reduce(threads_m: int, threads_k: int, threads_l):
m_tile = threads_m
k_tile = 128
mnk_tile = (m_tile, 1, k_tile)
@cute.kernel
def kernel(
mA_mkl: cute.Tensor,
mB_nkl: cute.Tensor,
mSFA_mkl: cute.Tensor,
mSFB_nkl: cute.Tensor,
mC_mnl: cute.Tensor,
):
tidx, tidy, tidz = arch.thread_idx()
bidx, bidy, bidz = arch.block_idx()
l_block = bidz * threads_l + tidz
# Tile views
gA_mkl = cute.local_tile(mA_mkl, slice_(mnk_tile, (None, 0, None)), (None, None, None))
gSFA_mkl = cute.local_tile(mSFA_mkl, slice_(mnk_tile, (None, 0, None)), (None, None, None))
gB_nkl = cute.local_tile(mB_nkl, slice_(mnk_tile, (0, None, None)), (None, None, None))
gSFB_nkl = cute.local_tile(mSFB_nkl, slice_(mnk_tile, (0, None, None)), (None, None, None))
gC_mnl = cute.local_tile(mC_mnl, slice_(mnk_tile, (None, None, 0)), (None, None, None))
tAgA0 = gA_mkl[tidx, None, bidx, 0, l_block]
tAgSFA0 = gSFA_mkl[tidx, None, bidx, 0, l_block]
tABrAB = cute.make_rmem_tensor_like(tAgA0, c_dtype)
tSFrSF = cute.make_rmem_tensor_like(tAgSFA0, sfc_dtype)
tCgC = gC_mnl[tidx, None, bidx, bidy, l_block]
allocator = cutlass.utils.SmemAllocator()
layout = cute.make_layout((threads_m, threads_k, threads_l))
res = allocator.allocate_tensor(sfc_dtype, layout)
r0, r1 = sfc_dtype(0), sfc_dtype(0)
k_tile_cnt = cute.assume(gA_mkl.layout[3].shape, 64)
for k_block in range(tidy, k_tile_cnt, threads_k):
tAgA = gA_mkl[tidx, None, bidx, k_block, l_block]
tBgB = gB_nkl[0, None, bidy, k_block, l_block]
tAgSFA = gSFA_mkl[tidx, None, bidx, k_block, l_block]
tBgSFB = gSFB_nkl[0, None, bidy, k_block, l_block]
a_vec = tAgA.load().to(c_dtype)
b_vec = tBgB.load().to(c_dtype)
sfa_vec = TensorSSA(cvt_f8e4m3_f16_intr(tAgSFA.load(), k_tile), k_tile, sfc_dtype)
sfb_vec = TensorSSA(cvt_f8e4m3_f16_intr(tBgSFB.load(), k_tile), k_tile, sfc_dtype)
tABrAB.store(a_vec * b_vec)
tSFrSF.store(sfa_vec * sfb_vec)
for i in cutlass.range_constexpr(0, k_tile, 2):
r0, r1 = fma_f16x2(
(tABrAB[i], tABrAB[i + 1]),
(tSFrSF[i], tSFrSF[i + 1]),
(r0, r1),
)
res[tidx, tidy, tidz] = r0 + r1
arch.sync_threads()
if tidy == 0:
out = cute.zeros_like(tCgC, acc_dtype)
for i in cutlass.range_constexpr(threads_k):
out += res[tidx, i, tidz]
tCgC.store(out.to(c_dtype))
return
@cute.jit
def my_kernel(
a_ptr: Pointer,
b_ptr: Pointer,
sfa_ptr: Pointer,
sfb_ptr: Pointer,
c_ptr: Pointer,
size: tuple[int, int, int, int],
):
kernel(*make_tensors(a_ptr, b_ptr, sfa_ptr, sfb_ptr, c_ptr, size)).launch(
grid=(ceil_div(size[0], threads_m), 1, ceil_div(size[3], threads_l)),
block=[threads_m, threads_k, threads_l],
cluster=(1, 1, 1),
)
return
return my_kernel
def warp_shuffle_f32(threads_m: int, threads_k: int, threads_l: int):
m_tile = threads_m
k_tile = 128
mnk_tile = (m_tile, 1, k_tile)
@cute.kernel
def kernel(
mA_mkl: cute.Tensor,
mB_nkl: cute.Tensor,
mSFA_mkl: cute.Tensor,
mSFB_nkl: cute.Tensor,
mC_mnl: cute.Tensor,
):
tidy, tidx, tidz = arch.thread_idx()
bidx, bidy, bidz = arch.block_idx()
l_block = bidz * threads_l + tidz
# Tile views
gA_mkl = cute.local_tile(mA_mkl, slice_(mnk_tile, (None, 0, None)), (None, None, None))
gSFA_mkl = cute.local_tile(mSFA_mkl, slice_(mnk_tile, (None, 0, None)), (None, None, None))
gB_nkl = cute.local_tile(mB_nkl, slice_(mnk_tile, (0, None, None)), (None, None, None))
gSFB_nkl = cute.local_tile(mSFB_nkl, slice_(mnk_tile, (0, None, None)), (None, None, None))
gC_mnl = cute.local_tile(mC_mnl, slice_(mnk_tile, (None, None, 0)), (None, None, None))
# Per-(M,L) output tile
# Per-thread register tiles to hold A*B and SFA*SFB products
tAgA0 = gA_mkl[tidx, None, bidx, 0, l_block]
tAgSFA0 = gSFA_mkl[tidx, None, bidx, 0, l_block]
tABrAB = cute.make_rmem_tensor_like(tAgA0, c_dtype)
tSFrSF = cute.make_rmem_tensor_like(tAgSFA0, sfc_dtype)
tCgC = gC_mnl[tidx, None, bidx, bidy, l_block]
res = acc_dtype(0)
k_tile_cnt = cute.assume(gA_mkl.layout[3].shape, 64)
for k_block in range(tidy, k_tile_cnt, threads_k):
tAgA = gA_mkl[tidx, None, bidx, k_block, l_block]
tAgSFA = gSFA_mkl[tidx, None, bidx, k_block, l_block]
tBgB = gB_nkl[0, None, bidy, k_block, l_block]
tBgSFB = gSFB_nkl[0, None, bidy, k_block, l_block]
a_vec = tAgA.load().to(c_dtype)
b_vec = tBgB.load().to(c_dtype)
sfa_vec = TensorSSA(cvt_f8e4m3_f16_intr(tAgSFA.load(), k_tile), k_tile, sfc_dtype)
sfb_vec = TensorSSA(cvt_f8e4m3_f16_intr(tBgSFB.load(), k_tile), k_tile, sfc_dtype)
tABrAB.store(a_vec * b_vec)
tSFrSF.store(sfa_vec * sfb_vec)
r0, r1 = sfc_dtype(0), sfc_dtype(0)
for i in cutlass.range_constexpr(0, k_tile, 2):
r0, r1 = fma_f16x2(
(tABrAB[i], tABrAB[i + 1]),
(tSFrSF[i], tSFrSF[i + 1]),
(r0, r1),
)
res += r0 + r1
offset = threads_k >> 1
while offset > 0:
res += arch.shuffle_sync_bfly(res, offset, threads_k)
offset >>= 1
if tidy == 0:
out = scalar_to_ssa(res, acc_dtype)
tCgC.store(out.to(c_dtype))
return
@cute.jit
def my_kernel(
a_ptr: Pointer,
b_ptr: Pointer,
sfa_ptr: Pointer,
sfb_ptr: Pointer,
c_ptr: Pointer,
size: tuple[int, int, int, int],
):
kernel(*make_tensors(a_ptr, b_ptr, sfa_ptr, sfb_ptr, c_ptr, size)).launch(
grid=(ceil_div(size[0], threads_m), 1, ceil_div(size[3], threads_l)),
block=[threads_k, threads_m, threads_l],
cluster=(1, 1, 1),
)
return
return my_kernel
def warp_shuffle_f16(threads_m: int, threads_k: int, threads_l: int):
m_tile = threads_m
k_tile = 128
mnk_tile = (m_tile, 1, k_tile)
@cute.kernel
def kernel(
mA_mkl: cute.Tensor,
mB_nkl: cute.Tensor,
mSFA_mkl: cute.Tensor,
mSFB_nkl: cute.Tensor,
mC_mnl: cute.Tensor,
):
# block = [threads_k, threads_m, threads_l]
tidy, tidx, tidz = arch.thread_idx()
bidx, bidy, bidz = arch.block_idx()
l_block = bidz * threads_l + tidz
gA_mkl = cute.local_tile(mA_mkl, slice_(mnk_tile, (None, 0, None)), (None, None, None))
gSFA_mkl = cute.local_tile(mSFA_mkl, slice_(mnk_tile, (None, 0, None)), (None, None, None))
gB_nkl = cute.local_tile(mB_nkl, slice_(mnk_tile, (0, None, None)), (None, None, None))
gSFB_nkl = cute.local_tile(mSFB_nkl, slice_(mnk_tile, (0, None, None)), (None, None, None))
gC_mnl = cute.local_tile(mC_mnl, slice_(mnk_tile, (None, None, 0)), (None, None, None))
# Output tile: one scalar per (M, L) after reduction over tidx
tCgC = gC_mnl[tidx, None, bidx, bidy, l_block]
# Per-thread register tiles
tAgA0 = gA_mkl[tidx, None, bidx, 0, l_block]
tAgSFA0 = gSFA_mkl[tidx, None, bidx, 0, l_block]
tABrAB = cute.make_rmem_tensor_like(tAgA0, c_dtype)
tSFrSF = cute.make_rmem_tensor_like(tAgSFA0, sfc_dtype)
r0, r1 = sfc_dtype(0), sfc_dtype(0)
k_tile_cnt = cute.assume(gA_mkl.layout[3].shape, 64)
for k_block in range(tidy, k_tile_cnt, threads_k):
tAgA = gA_mkl[tidx, None, bidx, k_block, l_block]
tAgSFA = gSFA_mkl[tidx, None, bidx, k_block, l_block]
tBgB = gB_nkl[0, None, bidy, k_block, l_block]
tBgSFB = gSFB_nkl[0, None, bidy, k_block, l_block]
a_vec = tAgA.load().to(c_dtype)
b_vec = tBgB.load().to(c_dtype)
sfa_vec = TensorSSA(cvt_f8e4m3_f16_intr(tAgSFA.load(), k_tile), k_tile, sfc_dtype)
sfb_vec = TensorSSA(cvt_f8e4m3_f16_intr(tBgSFB.load(), k_tile), k_tile, sfc_dtype)
tABrAB.store(a_vec * b_vec)
tSFrSF.store(sfa_vec * sfb_vec)
for i in cutlass.range_constexpr(0, k_tile, 2):
r0, r1 = fma_f16x2(
(tABrAB[i], tABrAB[i + 1]),
(tSFrSF[i], tSFrSF[i + 1]),
(r0, r1),
)
res = r0 + r1
offset = threads_k >> 1
while offset > 0:
res += arch.shuffle_sync_bfly(res, offset, threads_k)
offset >>= 1
if tidy == 0:
out = scalar_to_ssa(res, sfc_dtype)
tCgC.store(out.to(c_dtype))
return
@cute.jit
def my_kernel(
a_ptr: Pointer,
b_ptr: Pointer,
sfa_ptr: Pointer,
sfb_ptr: Pointer,
c_ptr: Pointer,
size: tuple[int, int, int, int],
):
kernel(*make_tensors(a_ptr, b_ptr, sfa_ptr, sfb_ptr, c_ptr, size)).launch(
grid=(ceil_div(size[0], threads_m), 1, ceil_div(size[3], threads_l)),
block=[threads_k, threads_m, threads_l],
cluster=(1, 1, 1),
)
return
return my_kernel
class Shape(Enum):
First = 1
Second = 2
Third = 3
_compiled_kernels: dict[tuple[int, int, int], Shape] = {}
def compile_kernel(
threads_m: int,
threads_k: int,
threads_l: int,
shape: Shape,
size: tuple[int, int, int, int],
):
key = (threads_m, threads_k, threads_l)
if key in _compiled_kernels:
return _compiled_kernels[key]
my_kernel = (
warp_smem_reduce
if shape is Shape.First
else warp_shuffle_f32
if shape is Shape.Second
else warp_shuffle_f16
)
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)
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)
c_ptr = make_ptr(c_dtype, 0, cute.AddressSpace.gmem, assumed_align=16)
compiled = cute.compile(
my_kernel(threads_m, threads_k, threads_l),
a_ptr,
b_ptr,
sfa_ptr,
sfb_ptr,
c_ptr,
size,
options=[
cute.OptLevel(3),
cute.GPUArch("sm_100a"),
# cute.PtxasOptions("--maxrregcount=8"),
],
)
_compiled_kernels[key] = compiled
return compiled
def select_threads(m: int, k: int, l: int) -> tuple[int, int, int, Shape]:
if (m, k, l) == (7168, 16384, 1):
return (64, 16, 1, Shape.First)
if (m, k, l) == (4096, 7168, 8):
return (64, 4, 4, Shape.Second)
if (m, k, l) == (7168, 2048, 4):
return (256, 4, 1, Shape.Third)
return (128, 8, 1, Shape.Third)
def custom_kernel(data: input_t) -> output_t:
a, b, _, _, sfa_permuted, sfb_permuted, c = data
m, k_half, l = a.shape
k = k_half * 2
n = 1
threads_m, threads_k, threads_l, shape = select_threads(m, k, l)
compiled_func = compile_kernel(threads_m, threads_k, threads_l, shape, (m, n, k, l))
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)
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)
c_ptr = make_ptr(c_dtype, c.data_ptr(), cute.AddressSpace.gmem, assumed_align=16)
compiled_func(a_ptr, b_ptr, sfa_ptr, sfb_ptr, c_ptr, (m, n, k, l))
return c
scrolls · 730 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