Skip to content
KernelIndex
Search⌘K

gpt-5-2025-08-07 / triton8c14a2

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

Use it

Vendorable · source mirrored · Apache-2.0View source →

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

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

Benchmark evidence

17 measurements across 1 GPU, fastest first.

Operation / workload
Hardware
Latency
Rank
Observed
GEMM n256 k7168fp16 · [1, 7168]
NVIDIA B200
202.5µs
#4 of 7
2025-10-16
GEMM n256 k7168fp16 · [4, 7168]
NVIDIA B200
205.3µs
#4 of 7
2025-10-16
GEMM n256 k7168fp16 · [14, 7168]
NVIDIA B200
211.0µs
#4 of 7
2025-10-16
GEMM n256 k7168fp16 · [15, 7168]
NVIDIA B200
212.5µs
#4 of 7
2025-10-16
GEMM n256 k7168fp16 · [16, 7168]
NVIDIA B200
212.8µs
#4 of 7
2025-10-16
GEMM n256 k7168fp16 · [32, 7168]
NVIDIA B200
215.7µs
#4 of 7
2025-10-16
GEMM n256 k7168fp16 · [901, 7168]
NVIDIA B200
221.8µs
#4 of 7
2025-10-16
GEMM n256 k7168fp16 · [54, 7168]
NVIDIA B200
228.3µs
#4 of 7
2025-10-16
GEMM n256 k7168fp16 · [58, 7168]
NVIDIA B200
230.0µs
#4 of 7
2025-10-16
GEMM n256 k7168fp16 · [53, 7168]
NVIDIA B200
237.0µs
#4 of 7
2025-10-16
Show all 17 measurements ›
GEMM n256 k7168fp16 · [57, 7168]
NVIDIA B200
238.8µs
#4 of 7
2025-10-16
GEMM n256 k7168fp16 · [56, 7168]
NVIDIA B200
239.5µs
#4 of 7
2025-10-16
GEMM n256 k7168fp16 · [55, 7168]
NVIDIA B200
239.9µs
#4 of 7
2025-10-16
GEMM n256 k7168fp16 · [63, 7168]
NVIDIA B200
240.8µs
#4 of 7
2025-10-16
GEMM n256 k7168fp16 · [80, 7168]
NVIDIA B200
242.9µs
#4 of 7
2025-10-16
GEMM n256 k7168fp16 · [11948, 7168]
NVIDIA B200
592.2µs
#5 of 7
2025-10-16
GEMM n256 k7168fp16 · [14104, 7168]
NVIDIA B200
594.3µs
#5 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:8b38711a03b00ced0e765466e7c49b64224b5fa0c92786a0d080e0364295062b
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.

autotune@triton.autotune(
mmaacc += tl.dot(a, b)
num-warps = 8triton.Config({'BLOCK_M': 128, 'BLOCK_N': 128, 'BLOCK_K': 64, 'GROUP_M': 8}, num_warps=8, num_stages=4),
stages = 4triton.Config({'BLOCK_M': 128, 'BLOCK_N': 128, 'BLOCK_K': 64, 'GROUP_M': 8}, num_warps=8, num_stages=4),

Kernel source

main.py178 lines
import torch
import triton
import triton.language as tl


# Autotuned GEMM kernel for:
#   A: [M, K] fp16
#   B: [N, K] fp16
#   C: [M, N] fp16
# Computes: C = A @ B.T
# Optimized for N=256, K=7168; M is variable.
@triton.autotune(
    configs=[
        triton.Config({'BLOCK_M': 128, 'BLOCK_N': 128, 'BLOCK_K': 64,  'GROUP_M': 8}, num_warps=8, num_stages=4),
        triton.Config({'BLOCK_M': 64,  'BLOCK_N': 256, 'BLOCK_K': 64,  'GROUP_M': 8}, num_warps=8, num_stages=4),
        triton.Config({'BLOCK_M': 128, 'BLOCK_N': 64,  'BLOCK_K': 128, 'GROUP_M': 8}, num_warps=8, num_stages=4),
        triton.Config({'BLOCK_M': 32,  'BLOCK_N': 256, 'BLOCK_K': 128, 'GROUP_M': 8}, num_warps=4, num_stages=5),
    ],
    key=['M'],  # N and K are constant for this op
)
@triton.jit
def _gemm_n256_k7168_kernel(
    A_ptr, B_ptr, C_ptr,
    M, N, K,
    stride_am, stride_ak,
    stride_bn, stride_bk,
    stride_cm, stride_cn,
    BLOCK_M: tl.constexpr,
    BLOCK_N: tl.constexpr,
    BLOCK_K: tl.constexpr,
    GROUP_M: tl.constexpr,
):
    # Program IDs with swizzled 1D launch to improve L2 hit-rate across M
    pid = tl.program_id(axis=0)
    num_pid_m = tl.cdiv(M, BLOCK_M)
    num_pid_n = tl.cdiv(N, BLOCK_N)
    num_pid_in_group = GROUP_M * num_pid_n
    group_id = pid // num_pid_in_group
    first_pid_m = group_id * GROUP_M
    group_size_m = tl.minimum(num_pid_m - first_pid_m, GROUP_M)
    pid_m = first_pid_m + (pid % group_size_m)
    pid_n = (pid % num_pid_in_group) // group_size_m

    # Offsets for this tile
    offs_m = pid_m * BLOCK_M + tl.arange(0, BLOCK_M)
    offs_n = pid_n * BLOCK_N + tl.arange(0, BLOCK_N)
    offs_k = tl.arange(0, BLOCK_K)

    # Pointers to tiles of A and B; note B is [N, K], we load [BK, BN] and use dot(A[BM,BK], B[BK,BN])
    A_tile_ptrs = A_ptr + (offs_m[:, None] * stride_am + offs_k[None, :] * stride_ak)
    B_tile_ptrs = B_ptr + (offs_n[None, :] * stride_bn + offs_k[:, None] * stride_bk)

    # Accumulator in FP32
    acc = tl.zeros((BLOCK_M, BLOCK_N), dtype=tl.float32)

    # Masks for bounds
    m_mask = offs_m < M
    n_mask = offs_n < N

    # K loop
    k_iter = 0
    while k_iter < K:
        k_mask = (k_iter + offs_k) < K
        a = tl.load(A_tile_ptrs, mask=m_mask[:, None] & k_mask[None, :], other=0.0)
        b = tl.load(B_tile_ptrs, mask=k_mask[:, None] & n_mask[None, :], other=0.0)
        acc += tl.dot(a, b)
        # Advance pointers along K
        A_tile_ptrs += BLOCK_K * stride_ak
        B_tile_ptrs += BLOCK_K * stride_bk
        k_iter += BLOCK_K

    # Write back in FP16
    C_tile_ptrs = C_ptr + (offs_m[:, None] * stride_cm + offs_n[None, :] * stride_cn)
    tl.store(C_tile_ptrs, acc.to(tl.float16), mask=m_mask[:, None] & n_mask[None, :])


def _validate_inputs(A: torch.Tensor, B: torch.Tensor):
    if A.ndim != 2 or B.ndim != 2:
        raise ValueError(f"Expected 2D tensors, got A.ndim={A.ndim}, B.ndim={B.ndim}")
    M, K_a = A.shape
    N_b, K_b = B.shape
    if N_b != 256:
        raise ValueError(f"N must be 256; got B.shape[0]={N_b}")
    if K_a != 7168 or K_b != 7168:
        raise ValueError(f"K must be 7168; got A.shape[1]={K_a}, B.shape[1]={K_b}")
    if A.dtype != torch.float16 or B.dtype != torch.float16:
        raise TypeError(f"Expected A and B to be torch.float16; got {A.dtype} and {B.dtype}")
    return M, 256, 7168


def _select_run_device(A: torch.Tensor, B: torch.Tensor):
    # If CUDA not available
    if not torch.cuda.is_available():
        if A.is_cuda or B.is_cuda:
            raise RuntimeError("CUDA is not available, but at least one input tensor is on CUDA.")
        return None  # CPU-only scenario

    # CUDA available
    if A.is_cuda and B.is_cuda:
        if A.device != B.device:
            raise ValueError(f"A and B must be on the same CUDA device; got {A.device} and {B.device}")
        return A.device

    # If only one is CUDA, use that device; if both CPU, use current CUDA device
    if A.is_cuda:
        return A.device
    if B.is_cuda:
        return B.device
    return torch.device('cuda', torch.cuda.current_device())


def _launch_triton(A: torch.Tensor, B: torch.Tensor) -> torch.Tensor:
    M, N, K = _validate_inputs(A, B)
    run_device = _select_run_device(A, B)

    # CPU-only fallback if CUDA not available
    if run_device is None:
        # Keep CPU semantics
        return torch.matmul(A, B.t())

    # Move inputs to the chosen CUDA device (non-blocking when possible)
    A_dev = A.to(device=run_device, non_blocking=True)
    B_dev = B.to(device=run_device, non_blocking=True)

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

    # Strides in elements
    stride_am, stride_ak = A_dev.stride()
    stride_bn, stride_bk = B_dev.stride()
    stride_cm, stride_cn = C_dev.stride()

    # Grid: 1D with grouping across M to improve cache locality on B200
    grid = lambda META: (
        triton.cdiv(M, META['BLOCK_M']) * triton.cdiv(N, META['BLOCK_N']),
    )

    # Launch kernel
    _gemm_n256_k7168_kernel[grid](
        A_dev, B_dev, C_dev,
        M, N, K,
        stride_am, stride_ak,
        stride_bn, stride_bk,
        stride_cm, stride_cn,
    )

    # Restore output device semantics:
    # - If both inputs were CPU originally: return CPU tensor
    # - Else if any input was CUDA originally: return on that CUDA device (first CUDA in A,B order)
    if (not A.is_cuda) and (not B.is_cuda):
        return C_dev.to(device=A.device, non_blocking=True)
    target_out = A.device if A.is_cuda else (B.device if B.is_cuda else run_device)
    if target_out != run_device:
        return C_dev.to(device=target_out, non_blocking=True)
    return C_dev


def run(*args, **kwargs):
    """
    Entry point matching the reference signature.
    Usage:
      - run(A, B)
      - run(A=A, B=B)
    This will:
      - Move CPU tensors to GPU if CUDA is available
      - Validate shapes/dtypes (M,256,7168; fp16)
      - Launch a Triton-optimized GEMM for B200
      - Return result on the original device (CPU if inputs were CPU-only; otherwise on the input CUDA device)
    """
    if len(args) >= 2:
        A, B = args[0], args[1]
    else:
        A = kwargs.get('A', None)
        B = kwargs.get('B', None)
    if A is None or B is None:
        raise ValueError("run requires tensors A and B, either as positional args or kwargs.")

    return _launch_triton(A, B)
scrolls · 178 lines total

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

Best evidence level for this revision: reported

JSON