Skip to content
KernelIndex
Search⌘K

gpt-5 / triton714ae0

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

Use it

Vendorable · source mirrored · Apache-2.0View source →

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

main.py
curl "https://kernelindex.com/api/v1/implementations/flashinfer-gpt-5-triton-714ae0?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
Fused add RMSNorm h7168bf16 · [7168] · batch_size=64
NVIDIA B200
12.9µs
#4 of 7
2025-10-16
Fused add RMSNorm h7168bf16 · [7168] · batch_size=18
NVIDIA B200
12.9µs
#4 of 7
2025-10-16
Fused add RMSNorm h7168bf16 · [7168] · batch_size=32
NVIDIA B200
13.6µs
#4 of 7
2025-10-16
Fused add RMSNorm h7168bf16 · [7168] · batch_size=7
NVIDIA B200
13.8µs
#4 of 7
2025-10-16
Fused add RMSNorm h7168bf16 · [7168] · batch_size=1
NVIDIA B200
14.4µs
#4 of 7
2025-10-16
Fused add RMSNorm h7168bf16 · [7168] · batch_size=539
NVIDIA B200
14.4µs
#3 of 7
2025-10-16
Fused add RMSNorm h7168bf16 · [7168] · batch_size=11949
NVIDIA B200
115.2µs
#3 of 7
2025-10-16
Fused add RMSNorm h7168bf16 · [7168] · batch_size=14521
NVIDIA B200
136.7µs
#3 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:767df550ec2489626c493f4196b597cddf70cccb66ee5e67524ccd39fdc0bf28
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.py123 lines
import torch
import triton
import triton.language as tl


@triton.jit
def fused_add_rmsnorm_h7168_kernel(
    hidden_ptr,      # *bfloat16
    residual_ptr,    # *bfloat16
    weight_ptr,      # *bfloat16
    out_ptr,         # *bfloat16
    batch_size,      # int32
    ld_hidden,       # int32
    ld_residual,     # int32
    ld_out,          # int32
    eps: tl.constexpr,
    H: tl.constexpr,
    BLOCK_SIZE: tl.constexpr,
):
    row = tl.program_id(axis=0)
    if row >= batch_size:
        return

    # Base pointers for this row
    hidden_row = hidden_ptr + row * ld_hidden
    residual_row = residual_ptr + row * ld_residual
    out_row = out_ptr + row * ld_out

    sum_sq = 0.0
    # Accumulate sum of squares
    for col_start in range(0, H, BLOCK_SIZE):
        cols = col_start + tl.arange(0, BLOCK_SIZE)
        mask = cols < H
        h = tl.load(hidden_row + cols, mask=mask, other=0.0)
        r = tl.load(residual_row + cols, mask=mask, other=0.0)
        x = tl.cast(h, tl.float32) + tl.cast(r, tl.float32)
        sum_sq += tl.sum(x * x, axis=0)

    mean = sum_sq / tl.cast(H, tl.float32)
    inv_rms = tl.rsqrt(mean + eps)

    # Normalize, scale by weight and store
    for col_start in range(0, H, BLOCK_SIZE):
        cols = col_start + tl.arange(0, BLOCK_SIZE)
        mask = cols < H
        h = tl.load(hidden_row + cols, mask=mask, other=0.0)
        r = tl.load(residual_row + cols, mask=mask, other=0.0)
        w = tl.load(weight_ptr + cols, mask=mask, other=0.0)
        x = tl.cast(h, tl.float32) + tl.cast(r, tl.float32)
        y = x * inv_rms * tl.cast(w, tl.float32)
        tl.store(out_row + cols, tl.cast(y, tl.bfloat16), mask=mask)


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

    if hidden_states.ndim != 2 or residual.ndim != 2:
        raise ValueError("hidden_states and residual must be 2D tensors of shape [batch_size, hidden_size]")
    if weight.ndim != 1:
        raise ValueError("weight must be 1D tensor of shape [hidden_size]")

    B, H = hidden_states.shape
    if H != 7168:
        raise AssertionError(f"hidden_size must be 7168, got {H}")
    if residual.shape != hidden_states.shape:
        raise ValueError(f"residual must have same shape as hidden_states, got {residual.shape} vs {hidden_states.shape}")
    if weight.shape[0] != H:
        raise ValueError(f"weight must have shape [{H}], got {tuple(weight.shape)}")

    # Device checks and management
    devices = {hidden_states.device.type, residual.device.type, weight.device.type}
    has_cuda_input = "cuda" in devices
    if has_cuda_input and not torch.cuda.is_available():
        raise RuntimeError("At least one input tensor is on CUDA, but CUDA is not available.")

    if not torch.cuda.is_available():
        raise RuntimeError("CUDA device is required to run the Triton kernel.")

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

    # Preserve original output device
    out_device = hidden_states.device

    # Move to target device and ensure dtype/layout
    def to_device_bf16_contig(t):
        return t.to(device=target_device, dtype=torch.bfloat16, non_blocking=True).contiguous()

    hs_dev = to_device_bf16_contig(hidden_states)
    rs_dev = to_device_bf16_contig(residual)
    w_dev = to_device_bf16_contig(weight)

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

    # Launch kernel
    BLOCK_SIZE = 1024  # tuned for H=7168
    grid = (B,)
    fused_add_rmsnorm_h7168_kernel[grid](
        hs_dev, rs_dev, w_dev, out_dev,
        B, H, H, H,
        eps=1e-6,
        H=7168,
        BLOCK_SIZE=BLOCK_SIZE,
        num_warps=8,
        num_stages=2,
    )

    # Move result back to original device if needed
    if out_device != target_device:
        out = out_dev.to(device=out_device, non_blocking=True)
    else:
        out = out_dev

    return out
scrolls · 123 lines total

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

Best evidence level for this revision: reported

JSON