Skip to content
KernelIndex
Search⌘K

gpt-5-2025-08-07 / triton82d3cf

gpt-5-2025-08-07_triton_82d3cf · gpt-5-2025-08-07 · triton · Apache-2.0

Use it

Vendorable · source mirrored · Apache-2.0View source →

No package. Vendor the mirrored source: 196 lines, Apache-2.0, pinned at da91508.

main.py
curl "https://kernelindex.com/api/v1/implementations/flashinfer-gpt-5-2025-08-07-triton-82d3cf?include=source"
interfacetriton
revisionda915083d4c7
symbolrun
pathmain.py
Compatibility
measured onNVIDIA B200
declared hardwareNVIDIA B200
architecturessm_100
dtypesfp16

Benchmark evidence

25 measurements across 1 GPU, fastest first.

Operation / workload
Hardware
Latency
Rank
Observed
GEMM n128 k2048fp16 · [93, 2048]
NVIDIA B200
15.8µs
#2 of 7
2025-10-16
GEMM n128 k2048fp16 · [128, 2048]
NVIDIA B200
16.0µs
#2 of 7
2025-10-16
GEMM n128 k2048fp16 · [25, 2048]
NVIDIA B200
16.3µs
#2 of 7
2025-10-16
GEMM n128 k2048fp16 · [32, 2048]
NVIDIA B200
16.3µs
#2 of 7
2025-10-16
GEMM n128 k2048fp16 · [6, 2048]
NVIDIA B200
16.4µs
#2 of 7
2025-10-16
GEMM n128 k2048fp16 · [5, 2048]
NVIDIA B200
16.6µs
#2 of 7
2025-10-16
GEMM n128 k2048fp16 · [63, 2048]
NVIDIA B200
16.7µs
#2 of 7
2025-10-16
GEMM n128 k2048fp16 · [1, 2048]
NVIDIA B200
16.9µs
#2 of 7
2025-10-16
GEMM n128 k2048fp16 · [492, 2048]
NVIDIA B200
16.9µs
#2 of 7
2025-10-16
GEMM n128 k2048fp16 · [2, 2048]
NVIDIA B200
16.9µs
#2 of 7
2025-10-16
Show all 25 measurements ›
GEMM n128 k2048fp16 · [17, 2048]
NVIDIA B200
17.0µs
#3 of 7
2025-10-16
GEMM n128 k2048fp16 · [172, 2048]
NVIDIA B200
17.0µs
#2 of 7
2025-10-16
GEMM n128 k2048fp16 · [289, 2048]
NVIDIA B200
17.2µs
#2 of 7
2025-10-16
GEMM n128 k2048fp16 · [16, 2048]
NVIDIA B200
17.3µs
#2 of 7
2025-10-16
GEMM n128 k2048fp16 · [64, 2048]
NVIDIA B200
17.3µs
#2 of 7
2025-10-16
GEMM n128 k2048fp16 · [4, 2048]
NVIDIA B200
17.3µs
#2 of 7
2025-10-16
GEMM n128 k2048fp16 · [8, 2048]
NVIDIA B200
17.4µs
#2 of 7
2025-10-16
GEMM n128 k2048fp16 · [952, 2048]
NVIDIA B200
19.0µs
#3 of 7
2025-10-16
GEMM n128 k2048fp16 · [8828, 2048]
NVIDIA B200
20.0µs
#2 of 7
2025-10-16
GEMM n128 k2048fp16 · [11006, 2048]
NVIDIA B200
22.0µs
#2 of 7
2025-10-16
GEMM n128 k2048fp16 · [12251, 2048]
NVIDIA B200
22.6µs
#2 of 7
2025-10-16
GEMM n128 k2048fp16 · [12853, 2048]
NVIDIA B200
22.6µs
#2 of 7
2025-10-16
GEMM n128 k2048fp16 · [14915, 2048]
NVIDIA B200
24.2µs
#2 of 7
2025-10-16
GEMM n128 k2048fp16 · [34, 2048]
NVIDIA B200
24.4µs
#4 of 7
2025-10-16
GEMM n128 k2048fp16 · [16294, 2048]
NVIDIA B200
24.9µs
#2 of 7
2025-10-16

Reported · How evidence levels are derived →

Source and license

sourcehttps://huggingface.co/datasets/flashinfer-ai/flashinfer-trace
commitda915083d4c7c5e61aa3005e3d17ae488e0fc71c
revision digestsha256:83de24e1d5e85c4974210fc44e8a0b8fcdeece984d871a7de4ea6afa9e48a188
license declaredApache-2.0
license concludedApache-2.0
authorsgpt-5-2025-08-07
imported2026-08-20

Techniques

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

autotunetriton.Config({"BLOCK_M": 128, "BLOCK_N": 128, "BLOCK_K": 64, "GROUP_M": 8}, num_warps=8, num_stages=5),
mmaacc += tl.dot(a, tl.trans(b))
num-warps = 8triton.Config({"BLOCK_M": 128, "BLOCK_N": 128, "BLOCK_K": 64, "GROUP_M": 8}, num_warps=8, num_stages=5),
stages = 5triton.Config({"BLOCK_M": 128, "BLOCK_N": 128, "BLOCK_K": 64, "GROUP_M": 8}, num_warps=8, num_stages=5),

Kernel source

main.py196 lines
import math
import torch
import triton
import triton.language as tl


# Autotuned GEMM kernel specialized for N=128, K=2048
# Computes: C[M, N] = A[M, K] @ B[N, K]^T
# A: [M, K] fp16, B: [N, K] fp16, C: [M, N] fp16
configs = [
    # High-throughput default tile for Blackwell/Hopper-class GPUs
    triton.Config({"BLOCK_M": 128, "BLOCK_N": 128, "BLOCK_K": 64, "GROUP_M": 8}, num_warps=8, num_stages=5),
    # Smaller M tile for small/irregular M to improve occupancy
    triton.Config({"BLOCK_M": 64, "BLOCK_N": 128, "BLOCK_K": 64, "GROUP_M": 8}, num_warps=4, num_stages=5),
    # Deeper K chunk for bandwidth-bound cases
    triton.Config({"BLOCK_M": 128, "BLOCK_N": 128, "BLOCK_K": 128, "GROUP_M": 4}, num_warps=8, num_stages=4),
]


@triton.autotune(configs=configs, key=["M"])
@triton.jit
def gemm_n128_k2048_kernel(
    A_ptr, B_ptr, C_ptr,
    M,                                   # runtime M
    stride_am, stride_ak,                # A strides
    stride_bn, stride_bk,                # B strides (N, K)
    stride_cm, stride_cn,                # C strides
    K: tl.constexpr,                     # K compile-time constant (2048)
    N: tl.constexpr,                     # N compile-time constant (128)
    BLOCK_M: tl.constexpr,
    BLOCK_N: tl.constexpr,
    BLOCK_K: tl.constexpr,
    GROUP_M: tl.constexpr,
):
    # Program ids for 2D launch grid
    pid_m = tl.program_id(0)
    pid_n = tl.program_id(1)

    # Offsets for M and N dimensions
    offs_m = pid_m * BLOCK_M + tl.arange(0, BLOCK_M)
    offs_n = pid_n * BLOCK_N + tl.arange(0, BLOCK_N)

    # Accumulator in fp32 for numerical stability
    acc = tl.zeros((BLOCK_M, BLOCK_N), dtype=tl.float32)

    # Base pointers for the first K-block
    a_ptrs = A_ptr + (offs_m[:, None] * stride_am + tl.arange(0, BLOCK_K)[None, :] * stride_ak)
    b_ptrs = B_ptr + (offs_n[:, None] * stride_bn + tl.arange(0, BLOCK_K)[None, :] * stride_bk)

    # Masks for M/N boundaries, broadcast across K chunk
    mask_m = offs_m[:, None] < M
    mask_n = offs_n[:, None] < N

    # K is guaranteed to be divisible by BLOCK_K for this problem (2048)
    tl.static_assert(BLOCK_N == 128, "Kernel specialized for N tiles of 128.")
    tl.static_assert((K % BLOCK_K) == 0, "K must be divisible by BLOCK_K.")
    # Hint to compiler for better vectorization/coalescing
    tl.max_contiguous(tl.arange(0, BLOCK_K), 64)

    # Main K loop
    for k0 in range(0, K, BLOCK_K):
        a = tl.load(a_ptrs, mask=mask_m, other=0.0)
        b = tl.load(b_ptrs, mask=mask_n, other=0.0)
        # Compute partial matmul: (BM, BK) x (BK, BN)
        acc += tl.dot(a, tl.trans(b))
        # Advance pointers along K
        a_ptrs += BLOCK_K * stride_ak
        b_ptrs += BLOCK_K * stride_bk

    # Write back result (convert to fp16)
    c_ptrs = C_ptr + (offs_m[:, None] * stride_cm + offs_n[None, :] * stride_cn)
    store_mask = (offs_m[:, None] < M) & (offs_n[None, :] < N)
    tl.store(c_ptrs, acc.to(tl.float16), mask=store_mask)


def _assert_and_prepare_inputs(A: torch.Tensor, B: torch.Tensor):
    # Validate dtypes
    if A.dtype != torch.float16 or B.dtype != torch.float16:
        raise TypeError("A and B must be torch.float16 (float16) tensors.")
    # Validate ranks
    if A.ndim != 2 or B.ndim != 2:
        raise ValueError("A and B must be 2D tensors: A[M, K], B[N, K].")
    M, K_a = A.shape
    N, K_b = B.shape
    if K_a != 2048 or K_b != 2048:
        raise ValueError(f"K must be 2048. Got A.shape[1]={K_a}, B.shape[1]={K_b}.")
    if N != 128:
        raise ValueError(f"N must be 128. Got B.shape[0]={N}.")
    return M, N, K_a


def _select_device_and_move(A: torch.Tensor, B: torch.Tensor):
    # Determine computation device
    a_dev = A.device
    b_dev = B.device
    cuda_available = torch.cuda.is_available()

    # If any tensor is already on CUDA, use that device
    if a_dev.type == "cuda" or b_dev.type == "cuda":
        if not cuda_available:
            raise RuntimeError("CUDA is not available but at least one input tensor is on CUDA.")
        target_device = a_dev if a_dev.type == "cuda" else b_dev
        A_dev = A.to(device=target_device, non_blocking=True)
        B_dev = B.to(device=target_device, non_blocking=True)
        return target_device, A_dev, B_dev
    # Both tensors on CPU
    if not cuda_available:
        raise RuntimeError("CUDA is not available; Triton kernel requires a CUDA-capable GPU.")
    target_device = torch.device("cuda", 0)
    A_dev = A.to(device=target_device, non_blocking=True)
    B_dev = B.to(device=target_device, non_blocking=True)
    return target_device, A_dev, B_dev


def _launch_kernel(A_dev: torch.Tensor, B_dev: torch.Tensor, M: int, N: int, K: int):
    # Ensure contiguity for optimal memory access
    if not A_dev.is_contiguous():
        A_dev = A_dev.contiguous()
    if not B_dev.is_contiguous():
        B_dev = B_dev.contiguous()

    # Allocate output on device
    C_dev = torch.empty((M, N), dtype=torch.float16, device=A_dev.device)

    # Compute grid
    def grid(meta):
        return (
            triton.cdiv(M, meta["BLOCK_M"]),
            triton.cdiv(N, meta["BLOCK_N"]),
        )

    gemm_n128_k2048_kernel[grid](
        A_dev, B_dev, C_dev,
        M,
        A_dev.stride(0), A_dev.stride(1),
        B_dev.stride(0), B_dev.stride(1),
        C_dev.stride(0), C_dev.stride(1),
        K=K,
        N=N,
    )
    return C_dev


def run(*args, **kwargs):
    """
    Entry point: C = run(A, B)

    - Computes C = A @ B.T for:
        A: [M, 2048] float16
        B: [128, 2048] float16
        C: [M, 128] float16
    - Optimized Triton kernel tuned for NVIDIA B200-class GPUs.
    - Handles device management: moves inputs to GPU, runs kernel, and returns result on the original device of A if both inputs shared the same original device; otherwise returns on the compute device.

    Args:
      A (torch.Tensor): shape [M, 2048], dtype float16
      B (torch.Tensor): shape [128, 2048], dtype float16

    Returns:
      torch.Tensor: shape [M, 128], dtype float16
    """
    # Extract inputs from args/kwargs
    if len(args) == 2:
        A, B = args
    else:
        if "A" in kwargs and "B" in kwargs:
            A = kwargs["A"]
            B = kwargs["B"]
        else:
            raise ValueError("run expects two tensors: A and B, either as positional or keyword arguments.")

    # Validate shapes/dtypes
    M, N, K = _assert_and_prepare_inputs(A, B)

    # Remember original devices
    orig_dev_A = A.device
    orig_dev_B = B.device

    # Move to appropriate device
    device, A_dev, B_dev = _select_device_and_move(A, B)

    # Launch Triton kernel
    C_dev = _launch_kernel(A_dev, B_dev, M, N, K)

    # Move result back
    # If both inputs were originally on the same device, return result to that device.
    # Otherwise, return on the compute (CUDA) device.
    if orig_dev_A == orig_dev_B:
        target_out_dev = orig_dev_A
    else:
        target_out_dev = device

    if target_out_dev.type == "cuda":
        return C_dev.to(device=target_out_dev, non_blocking=True)
    else:
        return C_dev.cpu()
scrolls · 196 lines total

Source code from FlashInfer-Bench (flashinfer-ai/flashinfer-trace) · Apache-2.0

Best evidence level for this revision: reported

JSON