Skip to content
KernelIndex
Search⌘K

gpt-5 / triton13f897

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

Use it

Vendorable · source mirrored · Apache-2.0View source →

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

main.py
curl "https://kernelindex.com/api/v1/implementations/flashinfer-gpt-5-triton-13f897?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 h512bf16 · [512] · batch_size=18
NVIDIA B200
6.19µs
#1 of 7
2025-10-16
RMSNorm h512bf16 · [512] · batch_size=64
NVIDIA B200
6.19µs
#2 of 7
2025-10-16
RMSNorm h512bf16 · [512] · batch_size=7
NVIDIA B200
6.22µs
#2 of 7
2025-10-16
RMSNorm h512bf16 · [512] · batch_size=539
NVIDIA B200
6.26µs
#2 of 7
2025-10-16
RMSNorm h512bf16 · [512] · batch_size=32
NVIDIA B200
6.36µs
#3 of 7
2025-10-16
RMSNorm h512bf16 · [512] · batch_size=1
NVIDIA B200
6.40µs
#2 of 7
2025-10-16
RMSNorm h512bf16 · [512] · batch_size=11949
NVIDIA B200
14.0µs
#2 of 7
2025-10-16
RMSNorm h512bf16 · [512] · batch_size=14521
NVIDIA B200
14.5µs
#1 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:7284a875d71b8a1d56296494d3f3d3eb5bf6a51e5d02fd48e12e0f3ba59a9707
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.py105 lines
import torch
import triton
import triton.language as tl


@triton.jit
def rmsnorm_h512_kernel(
    x_ptr, w_ptr, y_ptr,
    stride_xb, stride_xh,
    stride_yb, stride_yh,
    stride_w,
    B,
    H: tl.constexpr,
    EPS: tl.constexpr,
    BLOCK_SIZE: tl.constexpr,
):
    row = tl.program_id(axis=0)
    cols = tl.arange(0, BLOCK_SIZE)
    row_mask = row < B
    col_mask = cols < H
    mask = row_mask & col_mask

    x_row_ptrs = x_ptr + row * stride_xb + cols * stride_xh
    w_ptrs = w_ptr + cols * stride_w
    y_row_ptrs = y_ptr + row * stride_yb + cols * stride_yh

    x_bf16 = tl.load(x_row_ptrs, mask=mask, other=0.0)
    x = x_bf16.to(tl.float32)

    # Compute mean of squares in FP32
    sq = x * x
    mean_sq = tl.sum(sq, axis=0) / H
    inv_rms = 1.0 / tl.sqrt(mean_sq + EPS)

    w_bf16 = tl.load(w_ptrs, mask=col_mask, other=0.0)
    w = w_bf16.to(tl.float32)

    y = (x * inv_rms) * w
    y_bf16 = y.to(tl.bfloat16)
    tl.store(y_row_ptrs, y_bf16, mask=mask)


def run(hidden_states, weight):
    if not isinstance(hidden_states, torch.Tensor) or not isinstance(weight, torch.Tensor):
        raise TypeError("Inputs must be torch.Tensors")

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

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

    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}")

    hs_dev = hidden_states.device
    w_dev = weight.device

    # Determine target CUDA device
    target_cuda_device = None
    if hs_dev.type == "cuda":
        target_cuda_device = hs_dev
    elif w_dev.type == "cuda":
        target_cuda_device = w_dev
    else:
        if torch.cuda.is_available():
            target_cuda_device = torch.device("cuda")
        else:
            raise RuntimeError("CUDA is required to run this Triton kernel, but no CUDA device is available.")

    if target_cuda_device.type != "cuda":
        raise RuntimeError("Target device must be a CUDA device.")

    # Move inputs to target CUDA device if needed (without modifying originals)
    x_gpu = hidden_states.to(device=target_cuda_device, non_blocking=False)
    w_gpu = weight.to(device=target_cuda_device, non_blocking=False)

    # Prepare output on CUDA
    y_gpu = torch.empty_like(x_gpu, device=target_cuda_device)

    # Launch kernel
    grid = lambda meta: (batch_size,)
    rmsnorm_h512_kernel[grid](
        x_gpu, w_gpu, y_gpu,
        x_gpu.stride(0), x_gpu.stride(1),
        y_gpu.stride(0), y_gpu.stride(1),
        w_gpu.stride(0),
        batch_size,
        H=512,
        EPS=1e-6,
        BLOCK_SIZE=512,
        num_warps=4,
        num_stages=2,
    )

    # Move result back to original hidden_states device
    y_out = y_gpu.to(device=hs_dev, non_blocking=False)
    return y_out
scrolls · 105 lines total

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

Best evidence level for this revision: reported

JSON