Skip to content
KernelIndex
Search⌘K

gpt-5 / tritonbfd137

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

Use it

Vendorable · source mirrored · Apache-2.0View source →

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

main.py
curl "https://kernelindex.com/api/v1/implementations/flashinfer-gpt-5-triton-bfd137?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
RMSNorm h128bf16 · [128] · batch_size=256
NVIDIA B200
6.17µs
#1 of 9
2025-10-16
RMSNorm h128bf16 · [128] · batch_size=32
NVIDIA B200
6.18µs
#1 of 9
2025-10-16
RMSNorm h128bf16 · [128] · batch_size=4
NVIDIA B200
6.18µs
#3 of 9
2025-10-16
RMSNorm h128bf16 · [128] · batch_size=24
NVIDIA B200
6.18µs
#2 of 9
2025-10-16
RMSNorm h128bf16 · [128] · batch_size=316
NVIDIA B200
6.19µs
#3 of 9
2025-10-16
RMSNorm h128bf16 · [128] · batch_size=136
NVIDIA B200
6.19µs
#2 of 9
2025-10-16
RMSNorm h128bf16 · [128] · batch_size=192
NVIDIA B200
6.19µs
#2 of 9
2025-10-16
RMSNorm h128bf16 · [128] · batch_size=1088
NVIDIA B200
6.33µs
#4 of 9
2025-10-16
RMSNorm h128bf16 · [128] · batch_size=2528
NVIDIA B200
8.15µs
#5 of 9
2025-10-16
RMSNorm h128bf16 · [128] · batch_size=2048
NVIDIA B200
8.16µs
#6 of 9
2025-10-16
Show all 14 measurements ›
RMSNorm h128bf16 · [128] · batch_size=49532
NVIDIA B200
32.7µs
#5 of 9
2025-10-16
RMSNorm h128bf16 · [128] · batch_size=65016
NVIDIA B200
40.8µs
#5 of 9
2025-10-16
RMSNorm h128bf16 · [128] · batch_size=396256
NVIDIA B200
210.9µs
#4 of 9
2025-10-16
RMSNorm h128bf16 · [128] · batch_size=520128
NVIDIA B200
274.5µs
#3 of 9
2025-10-16

Reported · How evidence levels are derived →

Source and license

sourcehttps://huggingface.co/datasets/flashinfer-ai/flashinfer-trace
commitda915083d4c7c5e61aa3005e3d17ae488e0fc71c
revision digestsha256:c55770a64a41eee2a6c36faecb076cf854373b3f0f473dd77a38d3a72876088a
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.py126 lines
import torch
import triton
import triton.language as tl

EPS = 1e-6
HIDDEN_SIZE = 128


@triton.jit
def rmsnorm_h128_kernel(
    X_ptr, W_ptr, Y_ptr,
    stride_x_bs, stride_x_h,
    stride_w,
    stride_y_bs, stride_y_h,
    B: tl.int32,
    H: tl.constexpr,
    EPS: tl.constexpr,
):
    row = tl.program_id(0)
    offs = tl.arange(0, H)

    # Guards
    row_in_range = row < B
    vec_mask = row_in_range & (offs < H)

    x_ptrs = X_ptr + row * stride_x_bs + offs * stride_x_h
    w_ptrs = W_ptr + offs * stride_w
    y_ptrs = Y_ptr + row * stride_y_bs + offs * stride_y_h

    # Load inputs in bf16 then cast to f32
    x_bf16 = tl.load(x_ptrs, mask=vec_mask, other=0.0)
    x = x_bf16.to(tl.float32)

    # Compute RMS
    ss = tl.sum(x * x, axis=0)
    mean = ss * (1.0 / H)
    inv_rms = tl.rsqrt(mean + EPS)

    # Load weight
    w_bf16 = tl.load(w_ptrs, mask=offs < H, other=0.0)
    w = w_bf16.to(tl.float32)

    # Apply normalization and scale
    y = x * inv_rms * w
    y_bf16 = y.to(tl.bfloat16)

    # Store result
    tl.store(y_ptrs, y_bf16, mask=vec_mask)


def run(hidden_states, weight):
    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 != HIDDEN_SIZE:
        raise ValueError(f"hidden_size must be {HIDDEN_SIZE}, got {hidden_size}")
    if weight.shape[0] != HIDDEN_SIZE:
        raise ValueError(f"weight size must be [{HIDDEN_SIZE}], got {tuple(weight.shape)}")

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

    # Handle empty batch early
    if batch_size == 0:
        return hidden_states.new_empty((0, HIDDEN_SIZE))

    cuda_available = torch.cuda.is_available()
    orig_device = hidden_states.device

    # Device management
    if hidden_states.is_cuda and weight.is_cuda:
        if hidden_states.device != weight.device:
            target_device = hidden_states.device
            hidden_cuda = hidden_states
            weight_cuda = weight.to(target_device, non_blocking=True)
        else:
            target_device = hidden_states.device
            hidden_cuda = hidden_states
            weight_cuda = weight
    elif hidden_states.is_cuda and not weight.is_cuda:
        if not cuda_available:
            raise RuntimeError("CUDA is not available but hidden_states is on CUDA.")
        target_device = hidden_states.device
        hidden_cuda = hidden_states
        weight_cuda = weight.to(target_device, non_blocking=True)
    elif not hidden_states.is_cuda and weight.is_cuda:
        if not cuda_available:
            raise RuntimeError("CUDA is not available but weight is on CUDA.")
        target_device = weight.device
        hidden_cuda = hidden_states.to(target_device, non_blocking=True)
        weight_cuda = weight
    else:
        if not cuda_available:
            raise RuntimeError("CUDA is not available; Triton kernel requires CUDA.")
        target_device = torch.device("cuda")
        hidden_cuda = hidden_states.to(target_device, non_blocking=True)
        weight_cuda = weight.to(target_device, non_blocking=True)

    out_cuda = torch.empty((batch_size, HIDDEN_SIZE), dtype=torch.bfloat16, device=target_device)

    grid = (batch_size,)
    rmsnorm_h128_kernel[grid](
        hidden_cuda, weight_cuda, out_cuda,
        hidden_cuda.stride(0), hidden_cuda.stride(1),
        weight_cuda.stride(0),
        out_cuda.stride(0), out_cuda.stride(1),
        batch_size,
        H=HIDDEN_SIZE,
        EPS=EPS,
        num_warps=4,
        num_stages=2,
    )

    if out_cuda.device != orig_device:
        out = out_cuda.to(orig_device, non_blocking=True)
    else:
        out = out_cuda
    return out
scrolls · 126 lines total

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

Best evidence level for this revision: reported

JSON