Skip to content
KernelIndex
Search⌘K

gpt-o3 / triton38f281

gpt-o3_triton_38f281 · gpt-o3 · triton · Apache-2.0

Use it

Vendorable · source mirrored · Apache-2.0View source →

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

main.py
curl "https://kernelindex.com/api/v1/implementations/flashinfer-gpt-o3-triton-38f281?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
RMSNorm h2048bf16 · [2048] · batch_size=6
NVIDIA B200
12.3µs
#4 of 7
2025-10-16
RMSNorm h2048bf16 · [2048] · batch_size=79
NVIDIA B200
12.6µs
#5 of 7
2025-10-16
RMSNorm h2048bf16 · [2048] · batch_size=34
NVIDIA B200
12.6µs
#6 of 7
2025-10-16
RMSNorm h2048bf16 · [2048] · batch_size=64
NVIDIA B200
12.7µs
#6 of 7
2025-10-16
RMSNorm h2048bf16 · [2048] · batch_size=1
NVIDIA B200
14.3µs
#7 of 7
2025-10-16
RMSNorm h2048bf16 · [2048] · batch_size=12383
NVIDIA B200
80.5µs
#5 of 7
2025-10-16
RMSNorm h2048bf16 · [2048] · batch_size=16254
NVIDIA B200
102.3µs
#5 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:68d97ddb19595e999144e839d058825d4f1d00be60f48640ac7cee96bbf1da28
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 # good occupancy for B200
stages = 1num_stages=1,

Kernel source

main.py120 lines
import torch
import triton
import triton.language as tl

# ----------------------------------------------------------------------------- #
# Constants
# ----------------------------------------------------------------------------- #
HIDDEN_SIZE  = 2048           # fixed by specification
BLOCK_HIDDEN = 256            # elements processed per iteration
NUM_ITERS    = HIDDEN_SIZE // BLOCK_HIDDEN   # 8
NUM_WARPS    = 8              # good occupancy for B200

# ----------------------------------------------------------------------------- #
# Triton kernel
# ----------------------------------------------------------------------------- #
@triton.jit
def _rmsnorm_kernel(
    x_ptr,               # BF16 [B, H]
    w_ptr,               # BF16 [H]
    y_ptr,               # BF16 [B, H]
    stride_h,            # hidden size (2048)
    eps,                 # float, numerical stabiliser
    BLOCK_SIZE: tl.constexpr,
    N_ITERS:   tl.constexpr
):
    batch_id = tl.program_id(0)              # one program per row
    offs_h   = tl.arange(0, BLOCK_SIZE)      # [0, …, 255]

    row_start = batch_id * stride_h          # scalar

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

    for i in tl.static_range(N_ITERS):
        idx  = i * BLOCK_SIZE + offs_h
        x    = tl.load(x_ptr + row_start + idx).to(tl.float32)
        sum_sq += tl.sum(x * x, axis=0)

    mean_sq = sum_sq / stride_h
    inv_rms = tl.rsqrt(mean_sq + eps)        # scalar, broadcast later

    # ------------------------------------------------------------------ #
    # Pass 2 : normalise & scale
    # ------------------------------------------------------------------ #
    for i in tl.static_range(N_ITERS):
        idx = i * BLOCK_SIZE + offs_h

        x = tl.load(x_ptr + row_start + idx).to(tl.float32)
        w = tl.load(w_ptr + idx).to(tl.float32)

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


# ----------------------------------------------------------------------------- #
# Python wrapper
# ----------------------------------------------------------------------------- #
@torch.no_grad()
def run(hidden_states: torch.Tensor,
        weight:        torch.Tensor,
        *,
        eps: float = 1e-6) -> torch.Tensor:
    """
    RMSNorm (hidden_size = 2048) implemented with Triton.

    Args:
        hidden_states : [batch, 2048] bfloat16
        weight        : [2048]        bfloat16
        eps           : epsilon for numerical stability
    """
    # --------------- sanity checks --------------- #
    if hidden_states.dtype != torch.bfloat16 or weight.dtype != torch.bfloat16:
        raise TypeError("hidden_states and weight must be torch.bfloat16 tensors.")
    if hidden_states.ndim != 2 or hidden_states.shape[1] != HIDDEN_SIZE:
        raise ValueError(f"hidden_states must have shape [batch, {HIDDEN_SIZE}]")
    if weight.shape != (HIDDEN_SIZE,):
        raise ValueError(f"weight must have shape [{HIDDEN_SIZE}]")
    if not torch.cuda.is_available():
        raise RuntimeError("CUDA is required for Triton kernels.")

    orig_device = hidden_states.device

    xs_gpu = hidden_states.cuda(non_blocking=True) if not hidden_states.is_cuda else hidden_states
    w_gpu  = weight.cuda(non_blocking=True)        if not weight.is_cuda        else weight

    batch = xs_gpu.shape[0]
    y_gpu = torch.empty_like(xs_gpu)

    grid = (batch,)

    _rmsnorm_kernel[grid](
        xs_gpu,                       # x_ptr
        w_gpu,                        # w_ptr
        y_gpu,                        # y_ptr
        HIDDEN_SIZE,                  # stride_h
        eps,                          # eps
        BLOCK_HIDDEN,                 # BLOCK_SIZE
        NUM_ITERS,                    # N_ITERS
        num_warps=NUM_WARPS,
        num_stages=1,
    )

    return y_gpu.cpu() if orig_device.type == "cpu" else y_gpu


# ------------------------------ quick test ------------------------------ #
if __name__ == "__main__":
    torch.manual_seed(0)
    B = 4
    hs = torch.randn(B, HIDDEN_SIZE, dtype=torch.bfloat16)
    w  = torch.randn(HIDDEN_SIZE,     dtype=torch.bfloat16)

    ref = (hs.float() *
           torch.rsqrt(hs.float().pow(2).mean(-1, keepdim=True) + 1e-6) *
           w.float()).to(torch.bfloat16)

    out = run(hs, w)
    print("max error:", (ref - out).float().abs().max().item())
scrolls · 120 lines total

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

Best evidence level for this revision: reported

JSON