Skip to content
KernelIndex
Search⌘K

gpt-5 / triton0de5b5

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

Use it

Vendorable · source mirrored · Apache-2.0View source →

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

main.py
curl "https://kernelindex.com/api/v1/implementations/flashinfer-gpt-5-triton-0de5b5?include=source"
interfacetriton
revisionda915083d4c7
symbolrun
pathmain.py
Compatibility
measured onNVIDIA B200
declared hardwareNVIDIA B200
architecturessm_100
dtypesbf16

Benchmark evidence

14 measurements across 1 GPU, fastest first.

Operation / workload
Hardware
Latency
Rank
Observed
Fused add RMSNorm h4096bf16 · [4096] · batch_size=63
NVIDIA B200
20.9µs
#7 of 8
2025-10-16
Fused add RMSNorm h4096bf16 · [4096] · batch_size=15
NVIDIA B200
21.0µs
#7 of 8
2025-10-16
Fused add RMSNorm h4096bf16 · [4096] · batch_size=34
NVIDIA B200
21.2µs
#7 of 8
2025-10-16
Fused add RMSNorm h4096bf16 · [4096] · batch_size=79
NVIDIA B200
21.7µs
#7 of 8
2025-10-16
Fused add RMSNorm h4096bf16 · [4096] · batch_size=16
NVIDIA B200
21.9µs
#7 of 8
2025-10-16
Fused add RMSNorm h4096bf16 · [4096] · batch_size=64
NVIDIA B200
22.0µs
#7 of 8
2025-10-16
Fused add RMSNorm h4096bf16 · [4096] · batch_size=170
NVIDIA B200
22.2µs
#7 of 8
2025-10-16
Fused add RMSNorm h4096bf16 · [4096] · batch_size=7
NVIDIA B200
22.5µs
#7 of 8
2025-10-16
Fused add RMSNorm h4096bf16 · [4096] · batch_size=1
NVIDIA B200
23.5µs
#6 of 8
2025-10-16
Fused add RMSNorm h4096bf16 · [4096] · batch_size=8804
NVIDIA B200
132.7µs
#6 of 8
2025-10-16
Show all 14 measurements ›
Fused add RMSNorm h4096bf16 · [4096] · batch_size=10827
NVIDIA B200
159.4µs
#5 of 8
2025-10-16
Fused add RMSNorm h4096bf16 · [4096] · batch_size=11832
NVIDIA B200
171.1µs
#6 of 8
2025-10-16
Fused add RMSNorm h4096bf16 · [4096] · batch_size=14418
NVIDIA B200
205.1µs
#6 of 8
2025-10-16
Fused add RMSNorm h4096bf16 · [4096] · batch_size=14509
NVIDIA B200
205.8µs
#6 of 8
2025-10-16

Reported · How evidence levels are derived →

Source and license

sourcehttps://huggingface.co/datasets/flashinfer-ai/flashinfer-trace
commitda915083d4c7c5e61aa3005e3d17ae488e0fc71c
revision digestsha256:58aacca9df2b462724997435b7292d8eeb032761f7016688d6d1ce6ebd2092af
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.

num-warps = 8num_warps=8,
stages = 2num_stages=2,

Kernel source

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


@triton.jit
def fused_add_rmsnorm_h4096_kernel(
    hidden_ptr, residual_ptr, weight_ptr, output_ptr,
    M,  # number of rows (batch size)
    stride_hs_m, stride_hs_n,
    stride_res_m, stride_res_n,
    stride_out_m, stride_out_n,
    H: tl.constexpr,       # hidden size, must be 4096
    EPS: tl.constexpr,     # epsilon for numerical stability
    BLOCK_SIZE: tl.constexpr,
):
    tl.static_assert(H == 4096)
    pid = tl.program_id(0)
    row_in_bounds = pid < M

    cols = tl.arange(0, BLOCK_SIZE)

    # First pass: compute sum of squares across the row to get RMS
    sumsq = tl.zeros([1], dtype=tl.float32)
    for col_start in range(0, H, BLOCK_SIZE):
        off = col_start + cols
        mask = row_in_bounds & (off < H)
        hs = tl.load(hidden_ptr + pid * stride_hs_m + off * stride_hs_n, mask=mask, other=0).to(tl.float32)
        rs = tl.load(residual_ptr + pid * stride_res_m + off * stride_res_n, mask=mask, other=0).to(tl.float32)
        x = hs + rs
        sumsq += tl.sum(x * x, axis=0)

    mean_sq = sumsq / H
    inv_rms = tl.rsqrt(mean_sq + EPS)

    # Second pass: apply normalization and weight, then store
    for col_start in range(0, H, BLOCK_SIZE):
        off = col_start + cols
        mask = row_in_bounds & (off < H)
        hs = tl.load(hidden_ptr + pid * stride_hs_m + off * stride_hs_n, mask=mask, other=0).to(tl.float32)
        rs = tl.load(residual_ptr + pid * stride_res_m + off * stride_res_n, mask=mask, other=0).to(tl.float32)
        w = tl.load(weight_ptr + off, mask=(off < H), other=1.0).to(tl.float32)
        y = (hs + rs) * inv_rms * w
        tl.store(output_ptr + pid * stride_out_m + off * stride_out_n, y.to(tl.bfloat16), mask=mask)


def run(hidden_states, residual, weight):
    # Validate CUDA availability
    if not torch.cuda.is_available():
        raise RuntimeError("CUDA is required to run this Triton kernel, but no CUDA device is available.")

    # Validate inputs
    if not isinstance(hidden_states, torch.Tensor) or not isinstance(residual, torch.Tensor) or not isinstance(weight, torch.Tensor):
        raise TypeError("All inputs must be torch.Tensor instances.")

    if hidden_states.ndim != 2 or residual.ndim != 2 or weight.ndim != 1:
        raise ValueError("hidden_states and residual must be 2D tensors; weight must be a 1D tensor.")

    if hidden_states.shape != residual.shape:
        raise ValueError(f"hidden_states and residual must have the same shape, got {hidden_states.shape} vs {residual.shape}.")

    B, H = hidden_states.shape
    if H != 4096:
        raise ValueError(f"hidden_size must be 4096, got {H}.")
    if weight.shape[0] != H:
        raise ValueError(f"weight must have shape ({H},), got {weight.shape}.")

    # Determine target CUDA device
    target_device = None
    for t in (hidden_states, residual, weight):
        if t.is_cuda:
            target_device = t.device
            break
    if target_device is None:
        target_device = torch.device("cuda")

    # Original device of the main output (align with hidden_states)
    out_device = hidden_states.device

    # Move to GPU and ensure dtype is bfloat16 as specified
    hs_dev = hidden_states.to(device=target_device, dtype=torch.bfloat16, non_blocking=True)
    res_dev = residual.to(device=target_device, dtype=torch.bfloat16, non_blocking=True)
    w_dev = weight.to(device=target_device, dtype=torch.bfloat16, non_blocking=True)

    # Allocate output on GPU
    out_dev = torch.empty_like(hs_dev, device=target_device, dtype=torch.bfloat16)

    # Handle empty batch gracefully
    if B == 0:
        return out_dev.to(out_device)

    # Compute strides in elements
    hs_stride_m, hs_stride_n = hs_dev.stride()
    res_stride_m, res_stride_n = res_dev.stride()
    out_stride_m, out_stride_n = out_dev.stride()

    # Launch kernel
    grid = (B,)
    fused_add_rmsnorm_h4096_kernel[grid](
        hs_dev, res_dev, w_dev, out_dev,
        B,
        hs_stride_m, hs_stride_n,
        res_stride_m, res_stride_n,
        out_stride_m, out_stride_n,
        H=4096,
        EPS=1e-5,
        BLOCK_SIZE=256,
        num_warps=8,
        num_stages=2,
    )

    # Move result back to original device of hidden_states
    if out_device.type == "cuda" and out_device != target_device:
        return out_dev.to(out_device, non_blocking=True)
    elif out_device.type != "cuda":
        return out_dev.to(out_device)
    else:
        return out_dev
scrolls · 118 lines total

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

Best evidence level for this revision: reported

JSON