Skip to content
KernelIndex
Search⌘K

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
NVFP4 GEMVsuite of 3 cases
NVIDIA B200
30.6µs
#157 of 678
2025-11-25

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.

fp4NVFP4 GEMV Optimization v59: All shapes dual accumulator optimization
shared-memorysmem_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