Skip to content
KernelIndex
Search⌘K

submission 171244

yasmine3457 · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

sub.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-nvfp4-gemm-171244?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
44.9µs
#258 of 369
2025-12-17

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:d17363e8c20a689eb654358e1a91ee6ee913daf2cde979060606bfd8fd432018
license declaredunknown
license concludedunknown
authorsyasmine3457
imported2026-08-26

Techniques

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

fp4Optimized Block-scaled FP4 GEMM for NVIDIA B200

Kernel source

sub.py165 lines
import torch

def custom_kernel(data):
    """
    Optimized Block-scaled FP4 GEMM for NVIDIA B200
    
    Args:
        data: Tuple of (a, b, sfa, sfb, sfa_permuted, sfb_permuted, c)
    
    Returns:
        c: Output tensor
    """
    a, b, sfa, sfb, sfa_permuted, sfb_permuted, c = data
    return block_scaled_gemm_optimized(a, b, sfa, sfb, c)


def block_scaled_gemm_optimized(a, b, sfa, sfb, c):
    """
    Optimized GEMM with efficient scale factor conversion
    
    Key optimizations:
    1. Keep scale factors on GPU during conversion
    2. Vectorized blocking operations
    3. Minimize memory allocations
    """
    M, K_half, L = a.shape
    N = b.shape[0]
    
    def to_blocked_gpu(input_matrix):
        """GPU-based blocked conversion - no CPU transfer"""
        rows, cols = input_matrix.shape
        n_row_blocks = (rows + 127) // 128
        n_col_blocks = (cols + 3) // 4
        
        # Pad if necessary
        if rows % 128 != 0 or cols % 4 != 0:
            padded = torch.zeros((n_row_blocks * 128, n_col_blocks * 4), 
                               dtype=input_matrix.dtype, device=input_matrix.device)
            padded[:rows, :cols] = input_matrix
        else:
            padded = input_matrix
        
        # Reshape and permute - all on GPU
        blocks = padded.view(n_row_blocks, 128, n_col_blocks, 4).permute(0, 2, 1, 3)
        rearranged = blocks.reshape(-1, 4, 32, 4).transpose(1, 2).reshape(-1, 32, 16)
        return rearranged.flatten()
    
    # Process each batch
    for l_idx in range(L):
        # Convert scale factors on GPU
        scale_a = to_blocked_gpu(sfa[:, :, l_idx])
        scale_b = to_blocked_gpu(sfb[:, :, l_idx])
        
        result = torch._scaled_mm(
            a[:, :, l_idx],
            b[:, :, l_idx].transpose(0, 1),
            scale_a,
            scale_b,
            bias=None,
            out_dtype=torch.float16,
        )
        
        c[:, :, l_idx] = result
    
    return c


def block_scaled_gemm_fallback(a, b, sfa, sfb, c):
    """
    Fallback implementation using simple scale factors
    Only used if permuted versions aren't available
    """
    M, K_half, L = a.shape
    N = b.shape[0]
    
    def to_blocked(input_matrix):
        rows, cols = input_matrix.shape
        n_row_blocks = (rows + 127) // 128
        n_col_blocks = (cols + 3) // 4
        
        blocks = input_matrix.view(n_row_blocks, 128, n_col_blocks, 4).permute(0, 2, 1, 3)
        rearranged = blocks.reshape(-1, 4, 32, 4).transpose(1, 2).reshape(-1, 32, 16)
        return rearranged.flatten()
    
    for l_idx in range(L):
        scale_a = to_blocked(sfa[:, :, l_idx].cpu()).cuda()
        scale_b = to_blocked(sfb[:, :, l_idx].cpu()).cuda()
        
        result = torch._scaled_mm(
            a[:, :, l_idx],
            b[:, :, l_idx].transpose(0, 1),
            scale_a,
            scale_b,
            bias=None,
            out_dtype=torch.float16,
        )
        
        c[:, :, l_idx] = result
    
    return c


if __name__ == "__main__":
    print("Testing Optimized FP4 GEMM kernel...")
    
    # Test all benchmark dimensions
    test_cases = [
        (128, 7168, 16384, 1),
        (128, 4096, 7168, 1),
        (128, 7168, 2048, 1),
    ]
    
    for M, N, K, L in test_cases:
        print(f"\nTesting: M={M}, N={N}, K={K}, L={L}")
        
        torch.manual_seed(42)
        
        # Create FP4 tensors
        a = torch.randint(-128, 128, (L, M, K // 2), dtype=torch.int8, device="cuda").permute(1, 2, 0)
        b = torch.randint(-128, 128, (L, N, K // 2), dtype=torch.int8, device="cuda").permute(1, 2, 0)
        a = a.view(torch.float4_e2m1fn_x2)
        b = b.view(torch.float4_e2m1fn_x2)
        
        # Create scale factors
        sf_k = (K + 15) // 16
        sfa = torch.randint(0, 4, (L, M, sf_k), dtype=torch.int8, device='cuda').to(dtype=torch.float8_e4m3fn).permute(1, 2, 0)
        sfb = torch.randint(0, 4, (L, N, sf_k), dtype=torch.int8, device='cuda').to(dtype=torch.float8_e4m3fn).permute(1, 2, 0)
        
        # Create permuted versions
        def create_permuted(mn, sf_k, l):
            atom_m = (32, 4)
            atom_k = 4
            mma_shape = (l, (mn + 127) // 128, (sf_k + 3) // 4, atom_m[0], atom_m[1], atom_k)
            rand_int = torch.randint(0, 4, mma_shape, dtype=torch.int8, device='cuda')
            permuted = rand_int.to(dtype=torch.float8_e4m3fn).permute(3, 4, 1, 5, 2, 0)
            return permuted
        
        sfa_permuted = create_permuted(M, sf_k, L)
        sfb_permuted = create_permuted(N, sf_k, L)
        
        c = torch.zeros((L, M, N), dtype=torch.float16, device="cuda").permute(1, 2, 0)
        
        # Warm up
        data = (a, b, sfa, sfb, sfa_permuted, sfb_permuted, c)
        for _ in range(3):
            _ = custom_kernel(data)
        
        # Time it
        torch.cuda.synchronize()
        start = torch.cuda.Event(enable_timing=True)
        end = torch.cuda.Event(enable_timing=True)
        
        start.record()
        for _ in range(10):
            result = custom_kernel(data)
        end.record()
        
        torch.cuda.synchronize()
        elapsed_ms = start.elapsed_time(end) / 10
        
        print(f"  Average time: {elapsed_ms:.3f} ms")
        print(f"  Output shape: {result.shape}")
    
    print("\n✓ All tests completed!")
scrolls · 165 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