Skip to content
KernelIndex
Search⌘K

gpt-5 / triton679e13

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

Use it

Vendorable · source mirrored · Apache-2.0View source →

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

main.py
curl "https://kernelindex.com/api/v1/implementations/flashinfer-gpt-5-triton-679e13?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
Fused add RMSNorm h2048bf16 · [2048] · batch_size=34
NVIDIA B200
13.3µs
#8 of 8
2025-10-16
Fused add RMSNorm h2048bf16 · [2048] · batch_size=6
NVIDIA B200
13.4µs
#8 of 8
2025-10-16
Fused add RMSNorm h2048bf16 · [2048] · batch_size=79
NVIDIA B200
13.6µs
#8 of 8
2025-10-16
Fused add RMSNorm h2048bf16 · [2048] · batch_size=64
NVIDIA B200
14.3µs
#7 of 8
2025-10-16
Fused add RMSNorm h2048bf16 · [2048] · batch_size=1
NVIDIA B200
14.3µs
#6 of 8
2025-10-16
Fused add RMSNorm h2048bf16 · [2048] · batch_size=12383
NVIDIA B200
92.2µs
#8 of 8
2025-10-16
Fused add RMSNorm h2048bf16 · [2048] · batch_size=16254
NVIDIA B200
116.5µs
#8 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:93af4c1b8b13889962aa4b988589c7fecde417b3a7d9e7f7cdccf017f0796092
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.py157 lines
import torch
import triton
import triton.language as tl


# Constants
HIDDEN_SIZE = 2048
BLOCK_SIZE = 256
EPS = 1e-6


@triton.jit
def _fused_add_rmsnorm_h2048_kernel(
    hidden_ptr, residual_ptr, weight_ptr, out_ptr,
    B,
    hs_stride_row, hs_stride_col,
    rs_stride_row, rs_stride_col,
    out_stride_row, out_stride_col,
    w_stride,
    eps,
    H: tl.constexpr,
    BLOCK: tl.constexpr,
):
    row = tl.program_id(0)
    if row >= B:
        return

    sumsq = 0.0
    # First pass: compute sum of squares across the row
    for col_start in range(0, H, BLOCK):
        offs = tl.arange(0, BLOCK)
        cols = col_start + offs
        mask = cols < H

        hs_ptrs = hidden_ptr + row * hs_stride_row + cols * hs_stride_col
        rs_ptrs = residual_ptr + row * rs_stride_row + cols * rs_stride_col

        hs = tl.load(hs_ptrs, mask=mask, other=0).to(tl.float32)
        rs = tl.load(rs_ptrs, mask=mask, other=0).to(tl.float32)
        x = hs + rs

        sumsq += tl.sum(x * x, axis=0)

    mean = sumsq / H
    inv_rms = tl.rsqrt(mean + eps)

    # Second pass: normalize and scale by weight
    for col_start in range(0, H, BLOCK):
        offs = tl.arange(0, BLOCK)
        cols = col_start + offs
        mask = cols < H

        hs_ptrs = hidden_ptr + row * hs_stride_row + cols * hs_stride_col
        rs_ptrs = residual_ptr + row * rs_stride_row + cols * rs_stride_col
        w_ptrs = weight_ptr + cols * w_stride
        out_ptrs = out_ptr + row * out_stride_row + cols * out_stride_col

        hs = tl.load(hs_ptrs, mask=mask, other=0).to(tl.float32)
        rs = tl.load(rs_ptrs, mask=mask, other=0).to(tl.float32)
        w = tl.load(w_ptrs, mask=mask, other=0).to(tl.float32)

        y = (hs + rs) * inv_rms * w
        tl.store(out_ptrs, y.to(tl.bfloat16), mask=mask)


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

    if hidden_states.dtype != torch.bfloat16 or residual.dtype != torch.bfloat16 or weight.dtype != torch.bfloat16:
        raise TypeError("Tensors must be of dtype torch.bfloat16")

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

    if hidden_states.shape != residual.shape:
        raise ValueError("hidden_states and residual must have the same shape")

    B, H = hidden_states.shape
    if H != HIDDEN_SIZE:
        raise ValueError(f"hidden_size must be {HIDDEN_SIZE}, got {H}")

    if weight.numel() != HIDDEN_SIZE:
        raise ValueError(f"weight must have {HIDDEN_SIZE} elements, got {weight.numel()}")

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

    # Determine target CUDA device
    target_cuda_device = None
    if hs_dev.type == "cuda":
        target_cuda_device = hs_dev
    if rs_dev.type == "cuda":
        if target_cuda_device is None:
            target_cuda_device = rs_dev
        elif rs_dev != target_cuda_device:
            raise ValueError("All CUDA tensors must be on the same device")
    if w_dev.type == "cuda":
        if target_cuda_device is None:
            target_cuda_device = w_dev
        elif w_dev != target_cuda_device:
            raise ValueError("All CUDA tensors must be on the same device")

    if target_cuda_device is None:
        if not cuda_available:
            raise RuntimeError("CUDA is required but not available.")
        target_cuda_device = torch.device("cuda")
    else:
        if not cuda_available:
            raise RuntimeError("CUDA is not available but tensors are on CUDA.")

    # Move to CUDA if needed
    hs_cuda = hidden_states if hidden_states.device == target_cuda_device else hidden_states.to(device=target_cuda_device, non_blocking=True)
    rs_cuda = residual if residual.device == target_cuda_device else residual.to(device=target_cuda_device, non_blocking=True)
    w_cuda = weight if weight.device == target_cuda_device else weight.to(device=target_cuda_device, non_blocking=True)

    # Early return for empty batch
    if B == 0:
        out_empty = torch.empty_like(hidden_states)
        return out_empty

    # Allocate output on CUDA
    out_cuda = torch.empty_like(hs_cuda)

    # Kernel launch
    grid = (B,)

    _fused_add_rmsnorm_h2048_kernel[grid](
        hs_cuda,
        rs_cuda,
        w_cuda,
        out_cuda,
        B,
        hs_cuda.stride(0), hs_cuda.stride(1),
        rs_cuda.stride(0), rs_cuda.stride(1),
        out_cuda.stride(0), out_cuda.stride(1),
        w_cuda.stride(0),
        float(EPS),
        H=HIDDEN_SIZE,
        BLOCK=BLOCK_SIZE,
        num_warps=8,
        num_stages=2,
    )

    # Move result back to original device of hidden_states
    if hs_dev == out_cuda.device:
        return out_cuda
    else:
        return out_cuda.to(hs_dev, non_blocking=True)
scrolls · 157 lines total

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

Best evidence level for this revision: reported

JSON