Skip to content
KernelIndex
Search⌘K

gpt-5 / triton2f0daa

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

Use it

Vendorable · source mirrored · Apache-2.0View source →

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

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

Benchmark evidence

7 measurements across 1 GPU, fastest first.

Operation / workload
Hardware
Latency
Rank
Observed
RMSNorm h2048bf16 · [2048] · batch_size=79
NVIDIA B200
12.3µs
#4 of 7
2025-10-16
RMSNorm h2048bf16 · [2048] · batch_size=34
NVIDIA B200
12.3µs
#4 of 7
2025-10-16
RMSNorm h2048bf16 · [2048] · batch_size=64
NVIDIA B200
12.3µs
#4 of 7
2025-10-16
RMSNorm h2048bf16 · [2048] · batch_size=6
NVIDIA B200
13.7µs
#5 of 7
2025-10-16
RMSNorm h2048bf16 · [2048] · batch_size=1
NVIDIA B200
14.2µs
#5 of 7
2025-10-16
RMSNorm h2048bf16 · [2048] · batch_size=12383
NVIDIA B200
48.7µs
#4 of 7
2025-10-16
RMSNorm h2048bf16 · [2048] · batch_size=16254
NVIDIA B200
59.6µs
#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:5e3ecb227a5dbe9053849455646b2abbcba5da7e5478978b76b3ea6ea6535d7d
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 = 4num_warps = 4
stages = 2num_stages = 2

Kernel source

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


@triton.jit
def rmsnorm_h2048_kernel(
    x_ptr,  # *bf16
    w_ptr,  # *bf16
    y_ptr,  # *bf16
    stride_x,  # elements between rows in x
    stride_y,  # elements between rows in y
    batch_size,
    eps,
    BLOCK_SIZE: tl.constexpr,
    H: tl.constexpr,
):
    pid = tl.program_id(0)
    # Each program handles one row
    x_row_ptr = x_ptr + pid * stride_x
    y_row_ptr = y_ptr + pid * stride_y

    # Pass 1: accumulate sum of squares in fp32
    acc = tl.zeros((), dtype=tl.float32)
    for col in tl.static_range(0, H, BLOCK_SIZE):
        offs = col + tl.arange(0, BLOCK_SIZE)
        mask = offs < H
        x_bf16 = tl.load(x_row_ptr + offs, mask=mask, other=0)
        x = x_bf16.to(tl.float32)
        acc += tl.sum(x * x, axis=0)

    mean = acc / tl.full((), H, dtype=tl.float32)
    inv_rms = tl.rsqrt(mean + eps)

    # Pass 2: compute normalized output and apply weight
    for col in tl.static_range(0, H, BLOCK_SIZE):
        offs = col + tl.arange(0, BLOCK_SIZE)
        mask = offs < H
        x_bf16 = tl.load(x_row_ptr + offs, mask=mask, other=0)
        w_bf16 = tl.load(w_ptr + offs, mask=mask, other=0)
        x = x_bf16.to(tl.float32)
        w = w_bf16.to(tl.float32)
        y = (x * inv_rms) * w
        tl.store(y_row_ptr + offs, y.to(tl.bfloat16), mask=mask)


def run(hidden_states, weight):
    if hidden_states is None or weight is None:
        raise ValueError("hidden_states and weight must be provided")

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

    batch_size, hidden_size = hidden_states.shape
    if hidden_size != 2048:
        raise AssertionError(f"hidden_size must be 2048, got {hidden_size}")
    if weight.shape[0] != hidden_size:
        raise ValueError(f"weight length must match hidden_size={hidden_size}, got {weight.shape[0]}")

    if hidden_states.dtype != torch.bfloat16:
        raise TypeError(f"hidden_states must be torch.bfloat16, got {hidden_states.dtype}")
    if weight.dtype != torch.bfloat16:
        raise TypeError(f"weight must be torch.bfloat16, got {weight.dtype}")

    # Device management
    cuda_available = torch.cuda.is_available()
    hs_dev = hidden_states.device
    w_dev = weight.device

    if not cuda_available:
        raise RuntimeError("CUDA is required to run the Triton kernel but is not available.")

    # Choose target CUDA device
    if hs_dev.type == "cuda":
        target_device = hs_dev
    elif w_dev.type == "cuda":
        target_device = w_dev
    else:
        target_device = torch.device("cuda")

    # Move tensors to target_device if needed and ensure contiguous for optimal access
    x_gpu = hidden_states.to(device=target_device, dtype=torch.bfloat16, non_blocking=True).contiguous()
    w_gpu = weight.to(device=target_device, dtype=torch.bfloat16, non_blocking=True).contiguous()

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

    # Strides in elements (PyTorch strides are in elements)
    stride_x = x_gpu.stride(0)
    stride_y = y_gpu.stride(0)

    # Kernel launch
    grid = (batch_size,)
    EPS = 1e-6

    # Tunable meta-parameters for B200
    BLOCK_SIZE = 256  # 2048 / 256 = 8 steps, good balance for occupancy and bandwidth
    num_warps = 4
    num_stages = 2

    rmsnorm_h2048_kernel[grid](
        x_gpu,
        w_gpu,
        y_gpu,
        stride_x,
        stride_y,
        batch_size,
        EPS,
        BLOCK_SIZE=BLOCK_SIZE,
        H=hidden_size,
        num_warps=num_warps,
        num_stages=num_stages,
    )

    # Move result back to original device of hidden_states
    if hs_dev != target_device:
        return y_gpu.to(hs_dev, non_blocking=True)
    return y_gpu
scrolls · 120 lines total

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

Best evidence level for this revision: reported

JSON