Skip to content
KernelIndex
Search⌘K

gpt-5 / tritonb6738d

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

Use it

Vendorable · source mirrored · Apache-2.0View source →

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

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

Benchmark evidence

8 measurements across 1 GPU, fastest first.

Operation / workload
Hardware
Latency
Rank
Observed
RMSNorm h1536bf16 · [1536] · batch_size=32
NVIDIA B200
10.5µs
#5 of 8
2025-10-16
RMSNorm h1536bf16 · [1536] · batch_size=64
NVIDIA B200
10.7µs
#6 of 8
2025-10-16
RMSNorm h1536bf16 · [1536] · batch_size=18
NVIDIA B200
10.8µs
#7 of 8
2025-10-16
RMSNorm h1536bf16 · [1536] · batch_size=7
NVIDIA B200
11.7µs
#7 of 8
2025-10-16
RMSNorm h1536bf16 · [1536] · batch_size=539
NVIDIA B200
12.1µs
#7 of 8
2025-10-16
RMSNorm h1536bf16 · [1536] · batch_size=1
NVIDIA B200
12.2µs
#8 of 8
2025-10-16
RMSNorm h1536bf16 · [1536] · batch_size=11949
NVIDIA B200
61.1µs
#6 of 8
2025-10-16
RMSNorm h1536bf16 · [1536] · batch_size=14521
NVIDIA B200
71.7µ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:1cd50aed81df912a84720829a511e3757de146706f614eaa43952000be92221b
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, # tuned for B200
stages = 3num_stages=3,

Kernel source

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


@triton.jit
def rmsnorm_h1536_kernel(
    hidden_states_ptr,
    weight_ptr,
    output_ptr,
    batch_size,
    stride_hs_batch,
    stride_hs_hidden,
    stride_w,
    stride_out_batch,
    stride_out_hidden,
    eps,
):
    pid = tl.program_id(0)
    if pid >= batch_size:
        return

    # Base pointers for the current row
    hs_row_ptr = hidden_states_ptr + pid * stride_hs_batch
    out_row_ptr = output_ptr + pid * stride_out_batch

    # Pass 1: compute sum of squares across the hidden dimension (H=1536)
    sum_sq = tl.zeros((), dtype=tl.float32)
    for start in tl.static_range(0, 1536, 256):
        offs = start + tl.arange(0, 256)
        mask = offs < 1536
        x = tl.load(hs_row_ptr + offs * stride_hs_hidden, mask=mask, other=0).to(tl.float32)
        sum_sq += tl.sum(x * x, axis=0)

    mean_sq = sum_sq / 1536.0
    inv_rms = tl.rsqrt(mean_sq + eps)

    # Pass 2: normalize and scale by weight
    for start in tl.static_range(0, 1536, 256):
        offs = start + tl.arange(0, 256)
        mask = offs < 1536
        x = tl.load(hs_row_ptr + offs * stride_hs_hidden, mask=mask, other=0).to(tl.float32)
        w = tl.load(weight_ptr + offs * stride_w, mask=mask, other=0).to(tl.float32)
        y = x * inv_rms * w
        tl.store(out_row_ptr + offs * stride_out_hidden, y.to(tl.bfloat16), mask=mask)


def run(hidden_states, weight, eps=1e-6):
    # Validate inputs
    if not isinstance(hidden_states, torch.Tensor) or not isinstance(weight, torch.Tensor):
        raise TypeError("hidden_states and weight must be torch.Tensor")

    if hidden_states.ndim != 2:
        raise ValueError(f"hidden_states must be 2D [batch_size, hidden_size], got shape {tuple(hidden_states.shape)}")
    if weight.ndim != 1:
        raise ValueError(f"weight must be 1D [hidden_size], got shape {tuple(weight.shape)}")

    batch_size, hidden_size = hidden_states.shape
    if hidden_size != 1536:
        raise ValueError(f"hidden_size must be 1536, got {hidden_size}")
    if weight.numel() != hidden_size:
        raise ValueError(f"weight length must be {hidden_size}, got {weight.numel()}")

    cuda_available = torch.cuda.is_available()
    any_cuda_input = hidden_states.is_cuda or weight.is_cuda

    if any_cuda_input and not cuda_available:
        raise RuntimeError("CUDA is not available but GPU tensors were provided.")
    if not cuda_available:
        raise RuntimeError("CUDA is required to run Triton kernels, but no CUDA device is available.")

    # Determine execution device
    if hidden_states.is_cuda:
        exec_device = hidden_states.device
    elif weight.is_cuda:
        exec_device = weight.device
    else:
        exec_device = torch.device("cuda")

    # Move tensors to GPU and cast to bfloat16
    hs_gpu = hidden_states.to(device=exec_device, dtype=torch.bfloat16, copy=False)
    weight_gpu = weight.to(device=exec_device, dtype=torch.bfloat16, copy=False)

    # Allocate output
    out_gpu = torch.empty_like(hs_gpu, dtype=torch.bfloat16, device=exec_device)

    # Early exit for empty batch
    if batch_size == 0:
        return out_gpu.to(hidden_states.device) if hidden_states.device != exec_device else out_gpu

    # Strides in elements
    s_hs_batch, s_hs_hidden = hs_gpu.stride()
    (s_w,) = weight_gpu.stride()
    s_out_batch, s_out_hidden = out_gpu.stride()

    # Kernel launch configuration: one program per batch row
    grid = (batch_size,)

    rmsnorm_h1536_kernel[grid](
        hs_gpu,
        weight_gpu,
        out_gpu,
        batch_size,
        s_hs_batch,
        s_hs_hidden,
        s_w,
        s_out_batch,
        s_out_hidden,
        float(eps),
        num_warps=8,    # tuned for B200
        num_stages=3,
    )

    # Move result back to original device of hidden_states
    return out_gpu.to(hidden_states.device) if hidden_states.device != exec_device else out_gpu
scrolls · 116 lines total

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

Best evidence level for this revision: reported

JSON