Skip to content
KernelIndex
Search⌘K

gpt-5-2025-08-07 / triton9b01eb

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

Use it

Vendorable · source mirrored · Apache-2.0View source →

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

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

Benchmark evidence

29 measurements across 1 GPU, fastest first.

Operation / workload
Hardware
Latency
Rank
Observed
GEMM n2048 k4096fp16 · [2, 4096]
NVIDIA B200
143.4µs
#4 of 7
2025-10-16
GEMM n2048 k4096fp16 · [4, 4096]
NVIDIA B200
145.6µs
#4 of 7
2025-10-16
GEMM n2048 k4096fp16 · [1, 4096]
NVIDIA B200
145.7µs
#5 of 7
2025-10-16
GEMM n2048 k4096fp16 · [8, 4096]
NVIDIA B200
145.8µs
#4 of 7
2025-10-16
GEMM n2048 k4096fp16 · [5, 4096]
NVIDIA B200
145.9µs
#4 of 7
2025-10-16
GEMM n2048 k4096fp16 · [6, 4096]
NVIDIA B200
146.3µs
#4 of 7
2025-10-16
GEMM n2048 k4096fp16 · [15, 4096]
NVIDIA B200
147.5µs
#4 of 7
2025-10-16
GEMM n2048 k4096fp16 · [17, 4096]
NVIDIA B200
147.7µs
#4 of 7
2025-10-16
GEMM n2048 k4096fp16 · [32, 4096]
NVIDIA B200
147.8µs
#4 of 7
2025-10-16
GEMM n2048 k4096fp16 · [16, 4096]
NVIDIA B200
148.3µs
#4 of 7
2025-10-16
Show all 29 measurements ›
GEMM n2048 k4096fp16 · [25, 4096]
NVIDIA B200
149.1µs
#4 of 7
2025-10-16
GEMM n2048 k4096fp16 · [34, 4096]
NVIDIA B200
149.4µs
#4 of 7
2025-10-16
GEMM n2048 k4096fp16 · [289, 4096]
NVIDIA B200
149.5µs
#4 of 7
2025-10-16
GEMM n2048 k4096fp16 · [969, 4096]
NVIDIA B200
149.6µs
#4 of 7
2025-10-16
GEMM n2048 k4096fp16 · [172, 4096]
NVIDIA B200
149.7µs
#4 of 7
2025-10-16
GEMM n2048 k4096fp16 · [952, 4096]
NVIDIA B200
150.0µs
#4 of 7
2025-10-16
GEMM n2048 k4096fp16 · [492, 4096]
NVIDIA B200
151.0µs
#4 of 7
2025-10-16
GEMM n2048 k4096fp16 · [63, 4096]
NVIDIA B200
151.6µs
#4 of 7
2025-10-16
GEMM n2048 k4096fp16 · [64, 4096]
NVIDIA B200
151.6µs
#4 of 7
2025-10-16
GEMM n2048 k4096fp16 · [93, 4096]
NVIDIA B200
153.5µs
#4 of 7
2025-10-16
GEMM n2048 k4096fp16 · [128, 4096]
NVIDIA B200
160.1µs
#4 of 7
2025-10-16
GEMM n2048 k4096fp16 · [8828, 4096]
NVIDIA B200
945.5µs
#4 of 7
2025-10-16
GEMM n2048 k4096fp16 · [11006, 4096]
NVIDIA B200
1.20ms
#4 of 7
2025-10-16
GEMM n2048 k4096fp16 · [11938, 4096]
NVIDIA B200
1.41ms
#4 of 7
2025-10-16
GEMM n2048 k4096fp16 · [12251, 4096]
NVIDIA B200
1.41ms
#4 of 7
2025-10-16
GEMM n2048 k4096fp16 · [12853, 4096]
NVIDIA B200
1.62ms
#4 of 7
2025-10-16
GEMM n2048 k4096fp16 · [14915, 4096]
NVIDIA B200
1.67ms
#4 of 7
2025-10-16
GEMM n2048 k4096fp16 · [16294, 4096]
NVIDIA B200
1.69ms
#4 of 7
2025-10-16
GEMM n2048 k4096fp16 · [15813, 4096]
NVIDIA B200
1.69ms
#4 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:e233d827702d82b001ae4ce0cc006a334d047bf2a3a6bcedd79b168eb75cdc6c
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': 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.py138 lines
import torch
import triton
import triton.language as tl


@triton.autotune(
    configs=[
        triton.Config({'BLOCK_M': 128, 'BLOCK_N': 256, 'BLOCK_K': 64}, num_warps=8, num_stages=4),
        triton.Config({'BLOCK_M': 128, 'BLOCK_N': 128, 'BLOCK_K': 128}, num_warps=8, num_stages=4),
        triton.Config({'BLOCK_M': 64, 'BLOCK_N': 256, 'BLOCK_K': 64}, num_warps=4, num_stages=4),
        triton.Config({'BLOCK_M': 64, 'BLOCK_N': 128, 'BLOCK_K': 64}, num_warps=4, num_stages=3),
    ],
    key=['M'],
)
@triton.jit
def gemm_n2048_k4096_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,
):
    # Program ids for 2D launch grid
    pid_m = tl.program_id(axis=0)
    pid_n = tl.program_id(axis=1)

    # Offsets for the current block
    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 the first K-slice of A and B for this tile
    a_ptrs = A_ptr + offs_m[:, None] * stride_am + offs_k[None, :] * stride_ak
    b_ptrs = B_ptr + offs_n[None, :] * stride_bn + offs_k[:, None] * stride_bk

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

    # Loop over K dimension
    k = 0
    while k < K:
        a = tl.load(a_ptrs, mask=(offs_m[:, None] < M) & (k + offs_k[None, :] < K), other=0.0)
        b = tl.load(b_ptrs, mask=(offs_n[None, :] < N) & (k + offs_k[:, None] < K), other=0.0)
        acc += tl.dot(a, b)
        k += BLOCK_K
        a_ptrs += BLOCK_K * stride_ak
        b_ptrs += BLOCK_K * stride_bk

    # Write back to C (cast to FP16)
    c = acc.to(tl.float16)
    c_ptrs = C_ptr + offs_m[:, None] * stride_cm + offs_n[None, :] * stride_cn
    tl.store(c_ptrs, c, mask=(offs_m[:, None] < M) & (offs_n[None, :] < N))


def _validate_inputs(A: torch.Tensor, B: torch.Tensor):
    if not isinstance(A, torch.Tensor) or not isinstance(B, torch.Tensor):
        raise TypeError("Inputs A and B must be torch.Tensor instances.")
    if A.dtype != torch.float16 or B.dtype != torch.float16:
        raise TypeError(f"Inputs must be torch.float16. Got A.dtype={A.dtype}, B.dtype={B.dtype}.")
    if A.ndim != 2 or B.ndim != 2:
        raise ValueError(f"Inputs must be 2D tensors. Got A.ndim={A.ndim}, B.ndim={B.ndim}.")
    M, K_a = A.shape
    N_b, K_b = B.shape
    if K_a != 4096:
        raise ValueError(f"A must have shape [M, 4096]. Got {A.shape}.")
    if N_b != 2048 or K_b != 4096:
        raise ValueError(f"B must have shape [2048, 4096]. Got {B.shape}.")
    return M, 2048, 4096


def _compute_device_for_inputs(A: torch.Tensor, B: torch.Tensor):
    a_dev = A.device
    b_dev = B.device
    cuda_available = torch.cuda.is_available()

    # If any input is CUDA, ensure CUDA is available and devices match
    if a_dev.type == 'cuda' or b_dev.type == 'cuda':
        if not cuda_available:
            raise RuntimeError("CUDA tensor provided but CUDA is not available.")
        if a_dev.type == 'cuda' and b_dev.type == 'cuda' and a_dev != b_dev:
            raise ValueError("A and B must be on the same CUDA device.")
        return a_dev if a_dev.type == 'cuda' else b_dev

    # Both on CPU
    if not cuda_available:
        raise RuntimeError("CUDA is required for this Triton kernel, but no CUDA device is available.")
    # Use current CUDA device
    idx = torch.cuda.current_device()
    return torch.device(f"cuda:{idx}")


def run(*args, **kwargs):
    # Extract A and B from args/kwargs
    if len(args) >= 2:
        A, B = args[0], args[1]
    else:
        if 'A' not in kwargs or 'B' not in kwargs:
            raise ValueError("run requires tensors A and B either as positional or keyword arguments.")
        A, B = kwargs['A'], kwargs['B']

    M, N, K = _validate_inputs(A, B)
    compute_device = _compute_device_for_inputs(A, B)

    # Track original devices to restore output
    orig_a_dev = A.device
    orig_b_dev = B.device
    return_to_cpu = (orig_a_dev.type != 'cuda') and (orig_b_dev.type != 'cuda')

    # Move inputs to compute device and ensure contiguous layout for best performance
    with torch.cuda.device(compute_device.index if compute_device.index is not None else 0):
        A_dev = A.to(device=compute_device, non_blocking=True)
        B_dev = B.to(device=compute_device, non_blocking=True)
        # Contiguous for coalesced memory accesses
        if not A_dev.is_contiguous():
            A_dev = A_dev.contiguous()
        if not B_dev.is_contiguous():
            B_dev = B_dev.contiguous()

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

        # Launch kernel
        grid = lambda META: (triton.cdiv(M, META['BLOCK_M']), triton.cdiv(N, META['BLOCK_N']))
        gemm_n2048_k4096_kernel[grid](
            A_dev, B_dev, C_dev,
            M, N, K,
            A_dev.stride(0), A_dev.stride(1),
            B_dev.stride(0), B_dev.stride(1),
            C_dev.stride(0), C_dev.stride(1),
        )

        # Restore output to original device(s)
        if return_to_cpu:
            return C_dev.cpu()
        else:
            # Keep on GPU device where inputs lived
            return C_dev
scrolls · 138 lines total

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

Best evidence level for this revision: reported

JSON