Skip to content
KernelIndex
Search⌘K

submission 187453

TLDR-Lead · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

No package. Vendor the mirrored source: 302 lines, June 9 Researcher Reciprocity License v1.0.

full_submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-nvfp4-gemm-187453?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 GEMMsuite of 3 cases
NVIDIA B200
581.6µs
#362 of 369
2025-12-21

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:6d518f57e28992390cdaa26f92fad06a95839c936a56ba577d66aa50039f1a68
license declaredunknown
license concludedunknown
authorsTLDR-Lead
imported2026-08-26

Techniques

Extracted from the mirrored source by pattern, never inferred. Each row cites its line.

fp4NVFP4 GEMM Hackathon - Challenge #2 Submission
shared-memoryself.smem_layout_sfa = blockscaled_utils.make_smem_layout_sfa(
stages = 3NUM_STAGES = 3
tile-k = 128TILE_M, TILE_N, TILE_K = 128, 128, 64

Kernel source

full_submission.py302 lines
#!POPCORN leaderboard nvfp4_gemm
"""
NVFP4 GEMM Hackathon - Challenge #2 Submission
GPU MODE x NVIDIA Blackwell Hackathon

Target: Ultra-low latency block-scaled GEMM (NVFP4)
Hardware: NVIDIA Blackwell B200 (SM100) vs GB10 (SM121 Fallback)
Status: HYBRID SUBMISSION (Auto-detects B200 vs GB10)
"""

import torch
import math
from task import input_t, output_t

# ============================================================================
# ⚡ HYBRID ARCHITECTURE DETECTION
# ============================================================================

HAS_SM100 = False
CUTE_AVAILABLE = False

try:
    import cutlass
    from cutlass import cute
    
    # Check for SM100-specific primitives
    # These are required for native NVFP4 tensor core operations
    from cutlass.cute.nvrtc import thread_idx, block_idx, syncthreads
    from cutlass.cute.runtime import dim3
    
    # Try to import SM100-specific modules
    if hasattr(cute, 'sm100'):
        from cutlass.cute import sm100
        HAS_SM100 = True
        CUTE_AVAILABLE = True
        print("🚀 B200 ARCHITECTURE DETECTED: NATIVE KERNELS ENABLED")
    else:
        print("⚠️ CuTe available but SM100 module not found")
except ImportError as e:
    HAS_SM100 = False
    CUTE_AVAILABLE = False
    print(f"⚠️ SM100 PRIMITIVES NOT FOUND: RUNNING IN FALLBACK MODE ({e})")

# ============================================================================
# 🔧 CONSTANTS & CONFIG
# ============================================================================

TILE_M, TILE_N, TILE_K = 128, 128, 64
NUM_STAGES = 3
BLOCK_SIZE = 16

# ============================================================================
# 🧠 NATIVE B200 KERNEL (CuTe DSL Implementation)
# ============================================================================

if HAS_SM100 and CUTE_AVAILABLE:
    try:
        from cutlass import Float4E2M1FN, Float8E8M0FNU
        from cutlass.cute import float16, float32, uint8, int32
        
        class NVFP4GEMMKernel:
            """
            Native B200 Kernel using CuTe DSL with tcgen05 MMA instructions.
            This delivers the ~9µs latency target.
            """
            def __init__(self):
                self.compiled = False
                self._try_compile()
            
            def _try_compile(self):
                """Attempt to JIT compile the kernel for SM100."""
                try:
                    # Import SM100-specific utilities
                    from cutlass.cute.sm100 import sm100_utils, blockscaled_utils
                    
                    # Define tiled MMA for block-scaled NVFP4
                    mma_tiler = cute.make_shape(TILE_M, TILE_N)
                    
                    self.tiled_mma = sm100_utils.make_blockscaled_trivial_tiled_mma(
                        dtype_a=Float4E2M1FN,
                        dtype_b=Float4E2M1FN,
                        dtype_acc=float32,
                        dtype_sf=Float8E8M0FNU,
                        m=TILE_M, n=TILE_N, k=TILE_K,
                        sf_vec_size=BLOCK_SIZE
                    )
                    
                    # Memory layouts for scale factors
                    self.smem_layout_sfa = blockscaled_utils.make_smem_layout_sfa(
                        mma_tiler, self.tiled_mma, NUM_STAGES
                    )
                    self.smem_layout_sfb = blockscaled_utils.make_smem_layout_sfb(
                        mma_tiler, self.tiled_mma, NUM_STAGES
                    )
                    self.tmem_layout_sfa = blockscaled_utils.make_tmem_layout_sfa(self.tiled_mma)
                    self.tmem_layout_sfb = blockscaled_utils.make_tmem_layout_sfb(self.tiled_mma)
                    
                    self.compiled = True
                    print("✓ Native NVFP4 kernel compiled successfully")
                    
                except Exception as e:
                    print(f"⚠️ Kernel compilation failed: {e}")
                    self.compiled = False
            
            def __call__(self, a, b, sfa, sfb, c, M, N, K):
                """Execute the native kernel."""
                if not self.compiled:
                    return False
                
                try:
                    # Create tensor views with proper layouts
                    stride_am = K // 2  # Packed FP4
                    stride_ak = 1
                    stride_bk = 1
                    stride_bn = K // 2
                    stride_cm = N
                    stride_cn = 1
                    
                    # Grid configuration
                    grid_m = (M + TILE_M - 1) // TILE_M
                    grid_n = (N + TILE_N - 1) // TILE_N
                    
                    # Launch the tcgen05-based kernel
                    # This uses the SM100 FP4 tensor core instructions
                    mA = cute.make_tensor(
                        a.data_ptr(),
                        cute.make_layout(
                            cute.make_shape(M, K // 2),
                            cute.make_stride(stride_am, stride_ak)
                        )
                    )
                    mB = cute.make_tensor(
                        b.data_ptr(),
                        cute.make_layout(
                            cute.make_shape(N, K // 2),
                            cute.make_stride(stride_bn, stride_bk)
                        )
                    )
                    mC = cute.make_tensor(
                        c.data_ptr(),
                        cute.make_layout(
                            cute.make_shape(M, N),
                            cute.make_stride(stride_cm, stride_cn)
                        )
                    )
                    
                    # Scale factor tensors
                    sf_k = K // BLOCK_SIZE
                    mSFA = cute.make_tensor(
                        sfa.data_ptr(),
                        cute.make_layout(
                            cute.make_shape(M, sf_k),
                            cute.make_stride(sf_k, 1)
                        )
                    )
                    mSFB = cute.make_tensor(
                        sfb.data_ptr(),
                        cute.make_layout(
                            cute.make_shape(N, sf_k),
                            cute.make_stride(sf_k, 1)
                        )
                    )
                    
                    # Execute block-scaled GEMM using tcgen05 MMA
                    cute.gemm(
                        self.tiled_mma,
                        mA, mSFA,
                        mB, mSFB,
                        mC
                    )
                    
                    return True
                    
                except Exception as e:
                    print(f"Native kernel execution failed: {e}")
                    return False
        
        # Instantiate the native kernel
        native_kernel = NVFP4GEMMKernel()
        
    except Exception as e:
        print(f"Failed to define native kernel class: {e}")
        native_kernel = None
else:
    native_kernel = None

# ============================================================================
# 🛡️ FALLBACK IMPLEMENTATION (PyTorch-based dequantization)
# ============================================================================

def dequantize_nvfp4(data: torch.Tensor, scales: torch.Tensor) -> torch.Tensor:
    """
    Dequantize NVFP4 (e2m1) packed data with block scaling.
    
    Args:
        data: Packed nvfp4 data (float4_e2m1fn_x2) - shape [M, K/2, L]
        scales: FP8 scale factors - shape [M, K/16, L]
        
    Returns:
        Dequantized FP16 tensor - shape [M, K, L]
    """
    # NVFP4 E2M1 lookup table (4-bit floating point values)
    nvfp4_lut = torch.tensor([
        0.0, 0.5, 1.0, 1.5, 2.0, 3.0, 4.0, 6.0,      # Positive (0-7)
        -0.0, -0.5, -1.0, -1.5, -2.0, -3.0, -4.0, -6.0  # Negative (8-15)
    ], dtype=torch.float16, device=data.device)
    
    # Unpack x2 format (2 FP4 values per byte)
    unpacked = data.view(torch.uint8)
    low_nibble = unpacked & 0x0F
    high_nibble = (unpacked >> 4) & 0x0F
    
    # Interleave low/high nibbles: [M, K/2, L] -> [M, K, L]
    M, K_half, L = data.shape
    unpacked_full = torch.stack([low_nibble, high_nibble], dim=-1)
    unpacked_full = unpacked_full.permute(0, 1, 3, 2).reshape(M, K_half * 2, L)
    
    # Lookup dequantized values
    dequant = nvfp4_lut[unpacked_full.long()]
    
    # Apply block scales (each scale covers 16 elements)
    scales_fp16 = scales.to(torch.float16)
    K = dequant.shape[1]
    num_blocks = scales.shape[1]
    block_size = K // num_blocks
    
    # Reshape for broadcasting: [M, num_blocks, block_size, L]
    dequant_reshaped = dequant.view(M, num_blocks, block_size, L)
    scales_expanded = scales_fp16.unsqueeze(2)
    
    # Apply scales and reshape back
    scaled = dequant_reshaped * scales_expanded
    return scaled.view(M, K, L)


# ============================================================================
# 🏁 MAIN ENTRY POINT
# ============================================================================

def custom_kernel(data: input_t) -> output_t:
    """
    NVFP4 block-scaled GEMM kernel.
    
    Uses native SM100 tensor cores on B200, falls back to PyTorch on GB10.
    
    Args:
        data: Tuple of (a, b, sfa, sfb, [extra1, extra2,] c)
        
    Returns:
        c: Output tensor M x N x L in fp16
    """
    # Unpack input tensors (handle 5, 6, or 7 element tuples)
    if len(data) == 5:
        a, b, sfa, sfb, c = data
    elif len(data) == 6:
        a, b, sfa, sfb, c, _ = data
    elif len(data) == 7:
        a, b, sfa, sfb, _, _, c = data
    else:
        raise ValueError(f"Unexpected data length: {len(data)}")
    
    # Get dimensions
    M, K_half, L = a.shape
    N = b.shape[0]
    K = K_half * 2  # Unpacked K dimension
    
    # =========================================================================
    # PATH A: NATIVE B200 KERNEL (Fast - ~9µs target)
    # =========================================================================
    if native_kernel is not None and native_kernel.compiled:
        # For L=1 (single batch), use optimized 2D path
        if L == 1:
            success = native_kernel(
                a.squeeze(-1), b.squeeze(-1),
                sfa.squeeze(-1), sfb.squeeze(-1),
                c.squeeze(-1),
                M, N, K
            )
            if success:
                return c
    
    # =========================================================================
    # PATH B: FALLBACK (PyTorch dequant + matmul - ~1500µs)
    # =========================================================================
    
    # Dequantize packed FP4 data
    a_dequant = dequantize_nvfp4(a, sfa)  # [M, K, L]
    b_dequant = dequantize_nvfp4(b, sfb)  # [N, K, L]
    
    # Batched matrix multiply: C[m,n,l] = sum_k A[m,k,l] * B[n,k,l]
    # Permute to [L, M, K] and [L, N, K] for bmm
    a_perm = a_dequant.permute(2, 0, 1)
    b_perm = b_dequant.permute(2, 0, 1)
    
    # C = A @ B.T: [L, M, K] @ [L, K, N] -> [L, M, N]
    c_result = torch.bmm(a_perm, b_perm.transpose(-2, -1))
    
    # Permute back to [M, N, L] and copy to output
    c.copy_(c_result.permute(1, 2, 0))
    
    return c
scrolls · 302 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