Skip to content
KernelIndex
Search⌘K

gpt-o3 / triton793f87

gpt-o3_triton_793f87 · gpt-o3 · triton · Apache-2.0

Use it

Vendorable · source mirrored · Apache-2.0View source →

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

main.py
curl "https://kernelindex.com/api/v1/implementations/flashinfer-gpt-o3-triton-793f87?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 h4096bf16 · [4096] · batch_size=16
NVIDIA B200
10.1µs
#3 of 6
2025-10-16
RMSNorm h4096bf16 · [4096] · batch_size=15
NVIDIA B200
10.1µs
#4 of 6
2025-10-16
RMSNorm h4096bf16 · [4096] · batch_size=34
NVIDIA B200
10.2µs
#4 of 6
2025-10-16
RMSNorm h4096bf16 · [4096] · batch_size=79
NVIDIA B200
10.2µs
#3 of 6
2025-10-16
RMSNorm h4096bf16 · [4096] · batch_size=64
NVIDIA B200
10.2µs
#4 of 6
2025-10-16
RMSNorm h4096bf16 · [4096] · batch_size=63
NVIDIA B200
10.2µs
#4 of 6
2025-10-16
RMSNorm h4096bf16 · [4096] · batch_size=170
NVIDIA B200
10.2µs
#3 of 6
2025-10-16
RMSNorm h4096bf16 · [4096] · batch_size=7
NVIDIA B200
10.2µs
#4 of 6
2025-10-16
RMSNorm h4096bf16 · [4096] · batch_size=1
NVIDIA B200
10.2µs
#4 of 6
2025-10-16
RMSNorm h4096bf16 · [4096] · batch_size=8804
NVIDIA B200
45.0µs
#2 of 6
2025-10-16
Show all 14 measurements ›
RMSNorm h4096bf16 · [4096] · batch_size=10827
NVIDIA B200
53.1µs
#2 of 6
2025-10-16
RMSNorm h4096bf16 · [4096] · batch_size=11832
NVIDIA B200
56.6µs
#3 of 6
2025-10-16
RMSNorm h4096bf16 · [4096] · batch_size=14418
NVIDIA B200
67.2µs
#3 of 6
2025-10-16
RMSNorm h4096bf16 · [4096] · batch_size=14509
NVIDIA B200
67.5µs
#3 of 6
2025-10-16

Reported · How evidence levels are derived →

Source and license

sourcehttps://huggingface.co/datasets/flashinfer-ai/flashinfer-trace
commitda915083d4c7c5e61aa3005e3d17ae488e0fc71c
revision digestsha256:9aecf55f4e835ea63c3f587303dea2cd44adbdff5cb2cba9ec795826a3443ce2
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,
stages = 4num_stages=4,

Kernel source

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


@triton.jit
def _rmsnorm_kernel(
    hidden_ptr,          # *bf16  [batch_size, hidden_size]
    weight_ptr,          # *bf16  [hidden_size]
    out_ptr,             # *bf16  [batch_size, hidden_size]
    hidden_stride,       # int    stride between consecutive rows of hidden_ptr/out_ptr
    out_stride,          # int    stride between consecutive rows of out_ptr
    eps: tl.constexpr,   # float  numerical stability term
    hidden_size: tl.constexpr,
    BLOCK_SIZE: tl.constexpr,
):
    pid = tl.program_id(0)                       # program id == row id
    hidden_row_ptr = hidden_ptr + pid * hidden_stride
    out_row_ptr = out_ptr + pid * out_stride

    offs = tl.arange(0, BLOCK_SIZE)

    # --------------------------------------------------------------------- #
    # Pass 1 : compute sum of squares -> inverse RMS
    # --------------------------------------------------------------------- #
    rms_acc = tl.zeros([], dtype=tl.float32)

    for start in range(0, hidden_size, BLOCK_SIZE):
        idx = start + offs
        mask = idx < hidden_size
        x_bf16 = tl.load(hidden_row_ptr + idx, mask=mask, other=0)
        x_f32 = x_bf16.to(tl.float32)
        rms_acc += tl.sum(x_f32 * x_f32, axis=0)

    inv_rms = tl.math.rsqrt(rms_acc / hidden_size + eps)

    # --------------------------------------------------------------------- #
    # Pass 2 : normalize and scale
    # --------------------------------------------------------------------- #
    for start in range(0, hidden_size, BLOCK_SIZE):
        idx = start + offs
        mask = idx < hidden_size

        x_bf16 = tl.load(hidden_row_ptr + idx, mask=mask, other=0)
        w_bf16 = tl.load(weight_ptr + idx, mask=mask, other=0)

        x = x_bf16.to(tl.float32)
        w = w_bf16.to(tl.float32)

        y = x * inv_rms * w
        y_bf16 = y.to(tl.bfloat16)

        tl.store(out_row_ptr + idx, y_bf16, mask=mask)


def run(hidden_states: torch.Tensor, weight: torch.Tensor, *, eps: float = 1.0e-5):
    """
    RMSNorm (hidden_size = 4096) implemented in Triton.

    Arguments
    ---------
    hidden_states : (batch_size, 4096)  bfloat16
    weight        : (4096,)             bfloat16
    eps           : float, numerical stability term
    """
    # ---------------------------  Sanity checks  -------------------------- #
    if hidden_states.dim() != 2:
        raise ValueError("hidden_states must be 2-D [batch, hidden_size]")
    batch_size, hidden_size = hidden_states.shape
    if hidden_size != 4096:
        raise ValueError(f"hidden_size must be 4096, got {hidden_size}")
    if weight.dim() != 1 or weight.numel() != 4096:
        raise ValueError("weight must have shape (4096,)")
    if hidden_states.dtype != torch.bfloat16 or weight.dtype != torch.bfloat16:
        raise ValueError("hidden_states and weight must be torch.bfloat16")
    if not torch.cuda.is_available():
        raise RuntimeError("CUDA device is required to run Triton kernel")

    # --------------------------  Device handling  ------------------------- #
    orig_device = hidden_states.device
    hidden_cuda = hidden_states.cuda() if not hidden_states.is_cuda else hidden_states
    weight_cuda = weight.cuda() if not weight.is_cuda else weight

    hidden_cuda = hidden_cuda.contiguous()
    weight_cuda = weight_cuda.contiguous()

    output_cuda = torch.empty_like(hidden_cuda)

    hidden_stride = hidden_cuda.stride(0)
    out_stride = output_cuda.stride(0)

    # -----------------------------  Launch  ------------------------------- #
    BLOCK_SIZE = 1024
    grid = (batch_size,)

    _rmsnorm_kernel[grid](
        hidden_cuda,
        weight_cuda,
        output_cuda,
        hidden_stride,
        out_stride,
        eps,
        hidden_size=4096,
        BLOCK_SIZE=BLOCK_SIZE,
        num_warps=8,
        num_stages=4,
    )

    # ---------------------------  Return  --------------------------------- #
    return output_cuda if orig_device.type == "cuda" else output_cuda.cpu()
scrolls · 111 lines total

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

Best evidence level for this revision: reported

JSON