Skip to content
KernelIndex
Search⌘K

gpt-o3 / triton951f7e

gpt-o3_triton_951f7e · gpt-o3 · triton · Apache-2.0

Use it

Vendorable · source mirrored · Apache-2.0View source →

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

main.py
curl "https://kernelindex.com/api/v1/implementations/flashinfer-gpt-o3-triton-951f7e?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 h7168bf16 · [7168] · batch_size=64
NVIDIA B200
55.3µs
#7 of 7
2025-10-16
RMSNorm h7168bf16 · [7168] · batch_size=539
NVIDIA B200
55.3µs
#7 of 7
2025-10-16
RMSNorm h7168bf16 · [7168] · batch_size=18
NVIDIA B200
56.2µs
#7 of 7
2025-10-16
RMSNorm h7168bf16 · [7168] · batch_size=32
NVIDIA B200
56.6µs
#7 of 7
2025-10-16
RMSNorm h7168bf16 · [7168] · batch_size=7
NVIDIA B200
60.0µs
#7 of 7
2025-10-16
RMSNorm h7168bf16 · [7168] · batch_size=1
NVIDIA B200
62.4µs
#7 of 7
2025-10-16
RMSNorm h7168bf16 · [7168] · batch_size=11949
NVIDIA B200
280.9µs
#6 of 7
2025-10-16
RMSNorm h7168bf16 · [7168] · batch_size=14521
NVIDIA B200
329.9µs
#6 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:ecd18d4707ef95a5682b0d00da7a6517a89d7d3fded081f8c3881b15dae174e3
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 = 4num_warps=4, num_stages=4
stages = 4num_warps=4, num_stages=4

Kernel source

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

# ----------------------------------------------------------------------------- #
# Constants                                                                     #
# ----------------------------------------------------------------------------- #
HIDDEN_SIZE = 7168          # fixed, by specification
BLOCK_SIZE  = 128           # per-program processed columns

# ----------------------------------------------------------------------------- #
# Triton Kernel                                                                 #
# ----------------------------------------------------------------------------- #
@triton.jit
def _rmsnorm_h7168_kernel(
    x_ptr,            # pointer to input  [batch, hidden]
    w_ptr,            # pointer to weight [hidden]
    y_ptr,            # pointer to output [batch, hidden]
    eps,              # epsilon (float32)
    hidden_size: tl.constexpr,           # == 7168
    BLOCK: tl.constexpr                  # == BLOCK_SIZE
):
    """
    RMSNorm kernel for a single row (program = one batch element).
    Uses two passes over the hidden dimension:
      1. compute sum of squares  -> inv_rms
      2. write normalized output
    """
    pid  = tl.program_id(0)                       # program (=row) index
    offs = tl.arange(0, BLOCK)                    # vector of column offsets

    # --------------------------------------------------------------------- #
    # Pass 1: compute mean square & inv_rms                                 #
    # --------------------------------------------------------------------- #
    sum_sq = tl.zeros((), dtype=tl.float32)

    for start in range(0, hidden_size, BLOCK):
        idx   = start + offs
        mask  = idx < hidden_size
        x     = tl.load(x_ptr + pid * hidden_size + idx,
                        mask=mask, other=0.).to(tl.float32)
        sum_sq += tl.sum(x * x, axis=0)

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

    # --------------------------------------------------------------------- #
    # Pass 2: write out normalized values                                   #
    # --------------------------------------------------------------------- #
    for start in range(0, hidden_size, BLOCK):
        idx  = start + offs
        mask = idx < hidden_size

        x = tl.load(x_ptr + pid * hidden_size + idx,
                    mask=mask, other=0.).to(tl.float32)
        w = tl.load(w_ptr + idx, mask=mask, other=0.).to(tl.float32)

        y = x * inv_rms * w
        tl.store(y_ptr + pid * hidden_size + idx,
                 y.to(tl.bfloat16), mask=mask)


# ----------------------------------------------------------------------------- #
# Python wrapper                                                                #
# ----------------------------------------------------------------------------- #
@torch.no_grad()
def run(hidden_states: torch.Tensor,
        weight:        torch.Tensor,
        eps: float = 1e-6):
    """
    Executes RMSNorm on BF16 tensors using the above Triton kernel.

    Args:
        hidden_states : [batch, 7168] BF16 tensor
        weight        : [7168]        BF16 tensor
        eps           : numerical stability constant
    Returns:
        output        : same shape / dtype as hidden_states
    """
    # --------------------------------------------------------------------- #
    # Input validation                                                      #
    # --------------------------------------------------------------------- #
    if hidden_states.shape[-1] != HIDDEN_SIZE:
        raise ValueError(f"hidden_size must be {HIDDEN_SIZE}")
    if hidden_states.dtype != torch.bfloat16 or weight.dtype != torch.bfloat16:
        raise TypeError("Inputs must be torch.bfloat16")
    if not torch.cuda.is_available():
        raise RuntimeError("CUDA is required to run this kernel")

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

    batch_size = x_dev.shape[0]
    y_dev      = torch.empty_like(x_dev)

    # --------------------------------------------------------------------- #
    # Kernel launch                                                         #
    # --------------------------------------------------------------------- #
    grid = (batch_size,)
    _rmsnorm_h7168_kernel[grid](
        x_dev, w_dev, y_dev,
        eps,
        HIDDEN_SIZE,
        BLOCK_SIZE,
        num_warps=4, num_stages=4
    )

    # --------------------------------------------------------------------- #
    # Output to original device                                             #
    # --------------------------------------------------------------------- #
    return y_dev.cpu() if orig_device.type == 'cpu' else y_dev
scrolls · 116 lines total

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

Best evidence level for this revision: reported

JSON