Skip to content
KernelIndex
Search⌘K

submission 317513

Julius T · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

solution.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-nvfp4-dual-gemm-317513?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 dual GEMMsuite of 4 cases
NVIDIA B200
115.1µs
#418 of 420
2026-01-09

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:db205ef4ef3b43e10e51fb663f16505b5cebca3238cde0f96279704de17264b4
license declaredunknown
license concludedunknown
authorsJulius T
imported2026-08-26

Kernel source

solution.py131 lines
"""
Block-Scaled Dual GEMM with SiLU Activation for NVIDIA B200

Operation: C = SiLU(A @ B1.T) * (A @ B2.T)

Pure PyTorch implementation (no Triton, no JIT).
"""

import torch


def ceil_div(a: int, b: int) -> int:
    """Ceiling division."""
    return (a + b - 1) // b


def to_blocked(input_matrix: torch.Tensor) -> torch.Tensor:
    """
    Convert scale factor tensor to blocked format for cuBLAS.
    """
    rows = input_matrix.shape[0]
    cols = input_matrix.shape[1]

    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().contiguous()


def custom_kernel(data: tuple) -> torch.Tensor:
    """
    Dual GEMM with SiLU activation.

    C = SiLU(A @ B1.T) * (A @ B2.T)
    """
    # Unpack inputs
    a, b1, b2, sfa_cpu, sfb1_cpu, sfb2_cpu, _, _, _, c = data

    # Get dimensions
    m, n, l = c.shape

    device = a.device
    sfa = sfa_cpu.to(device=device, non_blocking=True) if sfa_cpu.device != device else sfa_cpu
    sfb1 = sfb1_cpu.to(device=device, non_blocking=True) if sfb1_cpu.device != device else sfb1_cpu
    sfb2 = sfb2_cpu.to(device=device, non_blocking=True) if sfb2_cpu.device != device else sfb2_cpu

    # Single-batch path (L=1)
    if l == 1:
        # Convert scale factors to blocked format
        scale_a = to_blocked(sfa[:, :, 0])
        scale_b1 = to_blocked(sfb1[:, :, 0])
        scale_b2 = to_blocked(sfb2[:, :, 0])

        b1_t = b1[:, :, 0].t()
        b2_t = b2[:, :, 0].t()

        # First GEMM: A @ B1.T
        gemm1 = torch._scaled_mm(
            a[:, :, 0], b1_t,
            scale_a, scale_b1,
            bias=None, out_dtype=torch.float32,
        )

        # Second GEMM: A @ B2.T
        gemm2 = torch._scaled_mm(
            a[:, :, 0], b2_t,
            scale_a, scale_b2,
            bias=None, out_dtype=torch.float32,
        )

        # Fused SiLU and multiply using PyTorch
        # SiLU(x) = x * sigmoid(x)
        result = (torch.nn.functional.silu(gemm1) * gemm2).to(torch.float16)

        # Reshape to M x N x 1
        return result.unsqueeze(2)

    # Multi-batch case (L > 1)
    gemm1_results = []
    gemm2_results = []

    for l_idx in range(l):
        # Convert scale factors to blocked format
        scale_a = to_blocked(sfa[:, :, l_idx])
        scale_b1 = to_blocked(sfb1[:, :, l_idx])
        scale_b2 = to_blocked(sfb2[:, :, l_idx])

        b1_t = b1[:, :, l_idx].t()
        b2_t = b2[:, :, l_idx].t()

        # First GEMM: A @ B1.T
        gemm1 = torch._scaled_mm(
            a[:, :, l_idx], b1_t,
            scale_a, scale_b1,
            bias=None, out_dtype=torch.float32,
        )
        gemm1_results.append(gemm1)

        # Second GEMM: A @ B2.T
        gemm2 = torch._scaled_mm(
            a[:, :, l_idx], b2_t,
            scale_a, scale_b2,
            bias=None, out_dtype=torch.float32,
        )
        gemm2_results.append(gemm2)

    # Stack results: list of (M, N) -> (L, M, N) -> (M, N, L)
    gemm1_stacked = torch.stack(gemm1_results, dim=0).permute(1, 2, 0)
    gemm2_stacked = torch.stack(gemm2_results, dim=0).permute(1, 2, 0)

    # Fused SiLU and multiply
    result = (torch.nn.functional.silu(gemm1_stacked) * gemm2_stacked).to(torch.float16)

    return result


def kernel(data: tuple) -> torch.Tensor:
    """
    Main entry point for the kernel.

    Args:
        data: Tuple of (a, b1, b2, sfa, sfb1, sfb2, sfa_perm, sfb1_perm, sfb2_perm, c)

    Returns:
        Output tensor C = SiLU(A @ B1.T) * (A @ B2.T) in float16
    """
    return custom_kernel(data)
scrolls · 131 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