Skip to content
KernelIndex
Search⌘K

gpt-5-2025-08-07 / tritonffc694

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

Use it

Vendorable · source mirrored · Apache-2.0View source →

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

main.py
curl "https://kernelindex.com/api/v1/implementations/flashinfer-gpt-5-2025-08-07-triton-ffc694?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 n5120 k2048fp16 · [16, 2048]
NVIDIA B200
20.4µs
#3 of 6
2025-10-16
GEMM n5120 k2048fp16 · [34, 2048]
NVIDIA B200
20.4µs
#3 of 6
2025-10-16
GEMM n5120 k2048fp16 · [4, 2048]
NVIDIA B200
20.4µs
#3 of 6
2025-10-16
GEMM n5120 k2048fp16 · [25, 2048]
NVIDIA B200
20.4µs
#3 of 6
2025-10-16
GEMM n5120 k2048fp16 · [2, 2048]
NVIDIA B200
20.4µs
#3 of 6
2025-10-16
GEMM n5120 k2048fp16 · [8, 2048]
NVIDIA B200
20.4µs
#3 of 6
2025-10-16
GEMM n5120 k2048fp16 · [63, 2048]
NVIDIA B200
20.4µs
#3 of 6
2025-10-16
GEMM n5120 k2048fp16 · [6, 2048]
NVIDIA B200
20.4µs
#3 of 6
2025-10-16
GEMM n5120 k2048fp16 · [5, 2048]
NVIDIA B200
20.4µs
#3 of 6
2025-10-16
GEMM n5120 k2048fp16 · [1, 2048]
NVIDIA B200
20.5µs
#2 of 6
2025-10-16
Show all 25 measurements ›
GEMM n5120 k2048fp16 · [32, 2048]
NVIDIA B200
20.5µs
#3 of 6
2025-10-16
GEMM n5120 k2048fp16 · [93, 2048]
NVIDIA B200
20.5µs
#3 of 6
2025-10-16
GEMM n5120 k2048fp16 · [64, 2048]
NVIDIA B200
20.5µs
#3 of 6
2025-10-16
GEMM n5120 k2048fp16 · [128, 2048]
NVIDIA B200
20.5µs
#3 of 6
2025-10-16
GEMM n5120 k2048fp16 · [17, 2048]
NVIDIA B200
20.6µs
#3 of 6
2025-10-16
GEMM n5120 k2048fp16 · [172, 2048]
NVIDIA B200
20.7µs
#3 of 6
2025-10-16
GEMM n5120 k2048fp16 · [289, 2048]
NVIDIA B200
22.3µs
#3 of 6
2025-10-16
GEMM n5120 k2048fp16 · [492, 2048]
NVIDIA B200
23.2µs
#3 of 6
2025-10-16
GEMM n5120 k2048fp16 · [952, 2048]
NVIDIA B200
38.9µs
#3 of 6
2025-10-16
GEMM n5120 k2048fp16 · [8828, 2048]
NVIDIA B200
188.3µs
#3 of 6
2025-10-16
GEMM n5120 k2048fp16 · [11006, 2048]
NVIDIA B200
231.8µs
#4 of 6
2025-10-16
GEMM n5120 k2048fp16 · [12251, 2048]
NVIDIA B200
256.6µs
#3 of 6
2025-10-16
GEMM n5120 k2048fp16 · [12853, 2048]
NVIDIA B200
281.2µs
#4 of 6
2025-10-16
GEMM n5120 k2048fp16 · [14915, 2048]
NVIDIA B200
324.7µs
#4 of 6
2025-10-16
GEMM n5120 k2048fp16 · [16294, 2048]
NVIDIA B200
362.6µs
#4 of 6
2025-10-16

Reported · How evidence levels are derived →

Source and license

sourcehttps://huggingface.co/datasets/flashinfer-ai/flashinfer-trace
commitda915083d4c7c5e61aa3005e3d17ae488e0fc71c
revision digestsha256:27e8fd799035035fd7df4cdd5e8b660505d38f90519afe31140d65a795cb8e39
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, out_dtype=tl.float32)
num-warps = 8triton.Config({"BLOCK_M": 128, "BLOCK_N": 256, "BLOCK_K": 64}, num_warps=8, num_stages=4),
stages = 4triton.Config({"BLOCK_M": 128, "BLOCK_N": 256, "BLOCK_K": 64}, num_warps=8, num_stages=4),

Kernel source

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


@triton.autotune(
    configs=[
        # Use power-of-two tile sizes to satisfy tl.arange power-of-two range requirement.
        triton.Config({"BLOCK_M": 128, "BLOCK_N": 256, "BLOCK_K": 64},  num_warps=8, num_stages=4),
        triton.Config({"BLOCK_M": 128, "BLOCK_N": 512, "BLOCK_K": 64},  num_warps=8, num_stages=4),
        triton.Config({"BLOCK_M": 64,  "BLOCK_N": 256, "BLOCK_K": 128}, num_warps=4, num_stages=4),
        triton.Config({"BLOCK_M": 256, "BLOCK_N": 128, "BLOCK_K": 64},  num_warps=8, num_stages=4),
        triton.Config({"BLOCK_M": 64,  "BLOCK_N": 128, "BLOCK_K": 64},  num_warps=4, num_stages=4),
        triton.Config({"BLOCK_M": 256, "BLOCK_N": 256, "BLOCK_K": 64},  num_warps=8, num_stages=5),
    ],
    key=["M"],  # Autotune over M; N=5120, K=2048 are fixed
)
@triton.jit
def gemm_n5120_k2048_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,
):
    """
    Compute C[M, N] = A[M, K] @ B[N, K]^T
    A: [M, K] row-major (stride_am, stride_ak)
    B: [N, K] row-major (stride_bn, stride_bk) but we read B^T tiles
    C: [M, N] row-major (stride_cm, stride_cn)
    """
    # 2D launch grid over (M-tiles, N-tiles)
    pid_m = tl.program_id(axis=0)
    pid_n = tl.program_id(axis=1)

    # Offsets for current 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)

    # Help the compiler with alignment assumptions
    tl.multiple_of(offs_m, 16)
    tl.multiple_of(offs_n, 16)
    tl.multiple_of(offs_k, 16)

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

    # Base pointers for the first K-slice
    a_ptrs = A_ptr + (offs_m[:, None] * stride_am + offs_k[None, :] * stride_ak)  # [BM, BK]
    # Load B as KxN by addressing B[n, k] -> B^T[k, n]
    b_ptrs = B_ptr + (offs_k[:, None] * stride_bk + offs_n[None, :] * stride_bn)  # [BK, BN]

    # Iterate along K dimension
    for k in range(0, K, BLOCK_K):
        a_mask = (offs_m[:, None] < M) & (offs_k[None, :] + k < K)
        b_mask = (offs_k[:, None] + k < K) & (offs_n[None, :] < N)

        a = tl.load(a_ptrs, mask=a_mask, other=0.0)
        b = tl.load(b_ptrs, mask=b_mask, other=0.0)

        # Tensor Core accelerated: fp16 x fp16 -> fp32 accumulation
        acc += tl.dot(a, b, out_dtype=tl.float32)

        # Advance pointers to next K block
        a_ptrs += BLOCK_K * stride_ak
        b_ptrs += BLOCK_K * stride_bk

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


def _assert_and_normalize_inputs(A: torch.Tensor, B: torch.Tensor):
    if A is None or B is None:
        raise ValueError("Expected tensors A and B, got None.")
    if A.ndim != 2 or B.ndim != 2:
        raise ValueError(f"Expected 2D tensors for A and B, got A.ndim={A.ndim}, B.ndim={B.ndim}")

    M, K_a = A.shape
    N_b, K_b = B.shape

    if K_a != 2048:
        raise ValueError(f"K dimension of A must be 2048, got {K_a}")
    if K_b != 2048:
        raise ValueError(f"K dimension of B (second dim) must be 2048, got {K_b}")
    if N_b != 5120:
        raise ValueError(f"N dimension of B (first dim) must be 5120, got {N_b}")

    # Convert dtypes if needed
    if A.dtype != torch.float16:
        A = A.to(torch.float16)
    if B.dtype != torch.float16:
        B = B.to(torch.float16)

    # Ensure contiguous layout (row-major) for efficient strided access
    if not A.is_contiguous():
        A = A.contiguous()
    if not B.is_contiguous():
        B = B.contiguous()

    return A, B


def _call_triton_gemm(A: torch.Tensor, B: torch.Tensor, *, stream: torch.cuda.Stream | None = None):
    """
    Internal: launch Triton kernel. The 'stream' argument is accepted for API
    compatibility but not passed to Triton (Triton uses the current stream).
    """
    # Shapes
    M, K = A.shape
    N = B.shape[0]  # 5120 by contract

    # Allocate output on same device as inputs (GPU)
    C = torch.empty((M, N), dtype=torch.float16, device=A.device)

    # Extract strides (in elements)
    stride_am, stride_ak = A.stride()
    stride_bn, stride_bk = B.stride()
    stride_cm, stride_cn = C.stride()

    # Grid: 2D grid over M-tiles and N-tiles
    def grid(meta):
        BM = meta["BLOCK_M"]
        BN = meta["BLOCK_N"]
        return (triton.cdiv(M, BM), triton.cdiv(N, BN))

    # Launch kernel; do NOT pass 'stream' kwarg to Triton
    gemm_n5120_k2048_kernel[grid](
        A, B, C,
        M, N, K,
        stride_am, stride_ak,
        stride_bn, stride_bk,
        stride_cm, stride_cn,
    )

    return C


def run(*args, **kwargs):
    """
    Entry point: C = run(A, B, stream=None)
    Computes C = A @ B.T for:
      - A: [M, 2048] float16
      - B: [5120, 2048] float16
      - C: [M, 5120] float16

    Device management:
      - If inputs are on CPU and CUDA is available, they are moved to GPU for the Triton kernel
      - If any input is on CUDA but CUDA is not available, raises a clear error
      - Result is moved back to the device of A (first input), preserving original device
      - If CUDA is not available and both inputs are CPU tensors, falls back to torch.matmul on CPU
      - Optional 'stream' (torch.cuda.Stream) sets the current stream for copies and compute
    """
    # Unpack inputs from args/kwargs
    if len(args) >= 2:
        A, B = args[0], args[1]
    else:
        A = kwargs.get("A", None)
        B = kwargs.get("B", None)

    # Optional CUDA stream
    stream = kwargs.get("stream", None)

    # Validate shapes/dtypes and ensure contiguous layout
    A, B = _assert_and_normalize_inputs(A, B)

    a_dev = A.device
    b_dev = B.device
    cuda_available = torch.cuda.is_available()

    # CPU-only path
    if not cuda_available:
        if A.is_cuda or B.is_cuda:
            raise RuntimeError("CUDA tensor provided but CUDA is not available.")
        return torch.matmul(A, B.T)

    # CUDA available: ensure tensors on CUDA, respecting the provided stream
    if stream is not None:
        if not isinstance(stream, torch.cuda.Stream):
            raise TypeError("stream must be a torch.cuda.Stream or None")
        with torch.cuda.stream(stream):
            A_cuda = A if A.is_cuda else A.cuda(non_blocking=True)
            B_cuda = B if B.is_cuda else B.cuda(non_blocking=True)
            C_cuda = _call_triton_gemm(A_cuda, B_cuda, stream=stream)  # stream is ignored by Triton
            # Move result back to match A's original device (if needed)
            if a_dev.type == "cuda":
                C_out = C_cuda
            else:
                C_out = C_cuda.to(a_dev, non_blocking=True)
    else:
        # Default stream
        A_cuda = A if A.is_cuda else A.cuda(non_blocking=True)
        B_cuda = B if B.is_cuda else B.cuda(non_blocking=True)
        C_cuda = _call_triton_gemm(A_cuda, B_cuda, stream=None)
        if a_dev.type == "cuda":
            C_out = C_cuda
        else:
            C_out = C_cuda.to(a_dev, non_blocking=True)

    return C_out


if __name__ == "__main__":
    # Simple sanity check
    torch.manual_seed(0)
    M = 512  # example M
    A = torch.randn((M, 2048), dtype=torch.float16)
    B = torch.randn((5120, 2048), dtype=torch.float16)
    C = run(A, B)
    ref = torch.matmul(A, B.T)
    max_diff = (C.cpu().float() - ref.float()).abs().max().item()
    print("Max abs diff:", max_diff)
scrolls · 216 lines total

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

Best evidence level for this revision: reported

JSON