Skip to content
KernelIndex
Search⌘K

gpt-o3 / tritone5fd98

gpt-o3_triton_e5fd98 · gpt-o3 · triton · Apache-2.0

Use it

Vendorable · source mirrored · Apache-2.0View source →

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

main.py
curl "https://kernelindex.com/api/v1/implementations/flashinfer-gpt-o3-triton-e5fd98?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
32.7µs
#7 of 7
2025-10-16
Fused add RMSNorm h7168bf16 · [7168] · batch_size=32
NVIDIA B200
32.9µs
#7 of 7
2025-10-16
Fused add RMSNorm h7168bf16 · [7168] · batch_size=18
NVIDIA B200
34.5µs
#7 of 7
2025-10-16
Fused add RMSNorm h7168bf16 · [7168] · batch_size=539
NVIDIA B200
35.2µs
#7 of 7
2025-10-16
Fused add RMSNorm h7168bf16 · [7168] · batch_size=7
NVIDIA B200
36.1µs
#7 of 7
2025-10-16
Fused add RMSNorm h7168bf16 · [7168] · batch_size=1
NVIDIA B200
36.3µs
#7 of 7
2025-10-16
Fused add RMSNorm h7168bf16 · [7168] · batch_size=11949
NVIDIA B200
502.5µs
#7 of 7
2025-10-16
Fused add RMSNorm h7168bf16 · [7168] · batch_size=14521
NVIDIA B200
603.5µs
#7 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:a23031c5b9cec36a07ecf730e0c60e814d3e5045f62941b1fa704e4ed9d5c775
license declaredApache-2.0
license concludedApache-2.0
authorsgpt-o3
imported2026-08-20

Techniques

Extracted from the mirrored source by pattern, never inferred. Each row cites its line.

num-warps = 8num_warps=8

Kernel source

main.py108 lines
import math
import torch
import triton
import triton.language as tl

# ----------------------------------------------------------------------------- #
# Constants
# ----------------------------------------------------------------------------- #
_HIDDEN_SIZE: int = 7168        # fixed hidden dimension
_BLOCK_SIZE: int = 256          # elements processed per kernel instance
_EPS_DEFAULT: float = 1e-6      # numerical stability

# ----------------------------------------------------------------------------- #
# Triton kernel
# ----------------------------------------------------------------------------- #
@triton.jit
def _fused_add_rmsnorm_h7168_kernel(
    hidden_ptr,                      # *bf16 [batch, 7168]
    residual_ptr,                    # *bf16 [batch, 7168]
    weight_ptr,                      # *bf16 [7168]
    out_ptr,                         # *bf16 [batch, 7168]
    eps,                             # float32
    hidden_size: tl.constexpr,       # 7168 (compile-time constant)
    BLOCK_SIZE: tl.constexpr         # 256  (compile-time constant)
):
    pid = tl.program_id(axis=0)                      # row / batch index
    row_start = pid * hidden_size                    # base pointer for the row
    offs = tl.arange(0, BLOCK_SIZE)                  # element indices inside the block

    # --------------------------------------------------------------------- #
    # Pass 1 : compute mean of squares
    # --------------------------------------------------------------------- #
    sum_sq = tl.zeros((), tl.float32)

    for offset in tl.static_range(0, hidden_size, BLOCK_SIZE):
        idx  = offset + offs
        mask = idx < hidden_size

        h = tl.load(hidden_ptr   + row_start + idx, mask=mask, other=0).to(tl.float32)
        r = tl.load(residual_ptr + row_start + idx, mask=mask, other=0).to(tl.float32)
        x = h + r
        sum_sq += tl.sum(x * x, axis=0)

    mean_sq = sum_sq / hidden_size
    inv_rms = tl.rsqrt(mean_sq + eps)

    # --------------------------------------------------------------------- #
    # Pass 2 : normalize, scale and store
    # --------------------------------------------------------------------- #
    for offset in tl.static_range(0, hidden_size, BLOCK_SIZE):
        idx  = offset + offs
        mask = idx < hidden_size

        h = tl.load(hidden_ptr   + row_start + idx, mask=mask, other=0).to(tl.float32)
        r = tl.load(residual_ptr + row_start + idx, mask=mask, other=0).to(tl.float32)
        w = tl.load(weight_ptr   + idx,          mask=mask, other=0).to(tl.float32)

        y = (h + r) * inv_rms * w
        tl.store(out_ptr + row_start + idx, y.to(tl.bfloat16), mask=mask)

# ----------------------------------------------------------------------------- #
# Wrapper
# ----------------------------------------------------------------------------- #
def run(hidden_states: torch.Tensor,
        residual:      torch.Tensor,
        weight:        torch.Tensor,
        eps:           float = _EPS_DEFAULT) -> torch.Tensor:
    """
    Fused Add + RMSNorm for hidden size 7168 (BF16) using Triton on B200 GPUs.
    """
    # ---------------------------- Sanity checks --------------------------- #
    if hidden_states.shape != residual.shape:
        raise ValueError("hidden_states and residual must have identical shapes")
    if hidden_states.shape[-1] != _HIDDEN_SIZE:
        raise ValueError(f"hidden_size must be {_HIDDEN_SIZE}, "
                         f"got {hidden_states.shape[-1]}")
    if not torch.cuda.is_available():
        raise RuntimeError("CUDA is not available — cannot execute Triton kernel")

    # ------------------------- Device management -------------------------- #
    orig_device = hidden_states.device                                      # remember caller's device
    device      = torch.device("cuda")                                      # execute on default CUDA

    hidden_gpu   = hidden_states.to(device, dtype=torch.bfloat16, non_blocking=True)
    residual_gpu = residual.to(device,      dtype=torch.bfloat16, non_blocking=True)
    weight_gpu   = weight.to(device,        dtype=torch.bfloat16, non_blocking=True)

    batch_size = hidden_gpu.shape[0]
    output_gpu = torch.empty_like(hidden_gpu, device=device, dtype=torch.bfloat16)

    # ------------------------- Launch Triton kernel ----------------------- #
    grid = (batch_size,)
    _fused_add_rmsnorm_h7168_kernel[grid](
        hidden_gpu,
        residual_gpu,
        weight_gpu,
        output_gpu,
        eps,
        hidden_size=_HIDDEN_SIZE,
        BLOCK_SIZE=_BLOCK_SIZE,
        num_warps=8
    )

    # ----------------------------- Return --------------------------------- #
    if orig_device.type == "cuda":
        return output_gpu.to(orig_device, non_blocking=True)

    return output_gpu.cpu()
scrolls · 108 lines total

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

Best evidence level for this revision: reported

JSON