Skip to content
KernelIndex
Search⌘K

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
NVFP4 GEMVsuite of 3 cases
NVIDIA B200
26.0µs
#112 of 678
2025-11-30

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-memorydef 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