Skip to content
KernelIndex
Search⌘K

gpt-o3 / triton35b90e

gpt-o3_triton_35b90e · gpt-o3 · triton · Apache-2.0

Use it

Vendorable · source mirrored · Apache-2.0View source →

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

main.py
curl "https://kernelindex.com/api/v1/implementations/flashinfer-gpt-o3-triton-35b90e?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 h128bf16 · [128] · batch_size=316
NVIDIA B200
6.30µs
#8 of 9
2025-10-16
RMSNorm h128bf16 · [128] · batch_size=136
NVIDIA B200
6.31µs
#8 of 9
2025-10-16
RMSNorm h128bf16 · [128] · batch_size=192
NVIDIA B200
6.31µs
#8 of 9
2025-10-16
RMSNorm h128bf16 · [128] · batch_size=24
NVIDIA B200
6.31µs
#7 of 9
2025-10-16
RMSNorm h128bf16 · [128] · batch_size=32
NVIDIA B200
6.31µs
#7 of 9
2025-10-16
RMSNorm h128bf16 · [128] · batch_size=256
NVIDIA B200
6.32µs
#8 of 9
2025-10-16
RMSNorm h128bf16 · [128] · batch_size=1088
NVIDIA B200
6.34µs
#5 of 9
2025-10-16
RMSNorm h128bf16 · [128] · batch_size=4
NVIDIA B200
6.34µs
#8 of 9
2025-10-16
RMSNorm h128bf16 · [128] · batch_size=2048
NVIDIA B200
7.82µs
#3 of 9
2025-10-16
RMSNorm h128bf16 · [128] · batch_size=2528
NVIDIA B200
8.01µs
#1 of 9
2025-10-16
Show all 14 measurements ›
RMSNorm h128bf16 · [128] · batch_size=49532
NVIDIA B200
32.5µs
#3 of 9
2025-10-16
RMSNorm h128bf16 · [128] · batch_size=65016
NVIDIA B200
40.1µs
#3 of 9
2025-10-16
RMSNorm h128bf16 · [128] · batch_size=396256
NVIDIA B200
210.7µs
#2 of 9
2025-10-16
RMSNorm h128bf16 · [128] · batch_size=520128
NVIDIA B200
274.6µs
#4 of 9
2025-10-16

Reported · How evidence levels are derived →

Source and license

sourcehttps://huggingface.co/datasets/flashinfer-ai/flashinfer-trace
commitda915083d4c7c5e61aa3005e3d17ae488e0fc71c
revision digestsha256:00524a07dd549ea8aa15ba3b1cc31a91f1b9274b43407b71448c733b7eeb5385
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

Kernel source

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

# -----------------------------------------------------------------------------#
# Constants
# -----------------------------------------------------------------------------#
EPS: float = 1e-6          # numerical stability
HIDDEN_SIZE: int = 128     # problem-specific constant

# -----------------------------------------------------------------------------#
# Triton Kernel
# -----------------------------------------------------------------------------#
@triton.jit
def _rmsnorm_kernel(
    x_ptr,                      # [batch, hidden]  (BF16)
    w_ptr,                      # [hidden]         (BF16)
    o_ptr,                      # [batch, hidden]  (BF16)
    stride_x,                   # leading dimension of x
    stride_o,                   # leading dimension of o
    eps: tl.constexpr,          # epsilon
    hidden: tl.constexpr        # hidden size (128)
):
    pid = tl.program_id(axis=0)                         # one program = one row
    offs = tl.arange(0, hidden)                         # [0 .. 127]
    mask = offs < hidden                                # always true, kept for safety

    # -------------------------------------------------------------------------#
    # Load input row and weight vector
    # -------------------------------------------------------------------------#
    x_row_ptr = x_ptr + pid * stride_x + offs
    w_ptrs    = w_ptr + offs
    x_bf16    = tl.load(x_row_ptr, mask=mask, other=0.0)
    w_bf16    = tl.load(w_ptrs,    mask=mask, other=0.0)

    x_f32 = x_bf16.to(tl.float32)
    w_f32 = w_bf16.to(tl.float32)

    # -------------------------------------------------------------------------#
    # RMS computation
    # -------------------------------------------------------------------------#
    rsq   = x_f32 * x_f32
    mean  = tl.sum(rsq) / hidden
    inv_r = tl.rsqrt(mean + eps)

    # -------------------------------------------------------------------------#
    # Final output:  y = (x * inv_rms) * weight
    # -------------------------------------------------------------------------#
    y_f32 = (x_f32 * inv_r) * w_f32
    y_bf16 = y_f32.to(tl.bfloat16)

    # -------------------------------------------------------------------------#
    # Store
    # -------------------------------------------------------------------------#
    o_row_ptr = o_ptr + pid * stride_o + offs
    tl.store(o_row_ptr, y_bf16, mask=mask)


# -----------------------------------------------------------------------------#
# Python Wrapper
# -----------------------------------------------------------------------------#
def run(*args, **kwargs):
    """
    Entry point.

    Parameters (positional or keyword):
      hidden_states: Tensor[batch, 128] (bfloat16)
      weight:        Tensor[128]        (bfloat16)

    Returns:
      output Tensor with same shape/dtype/device as `hidden_states`
    """
    # -------------------------------------------------------------------------#
    # Argument extraction
    # -------------------------------------------------------------------------#
    if len(args) + len(kwargs) < 2:
        raise TypeError("run() missing required arguments: 'hidden_states' and 'weight'")

    hidden_states = kwargs.pop('hidden_states') if 'hidden_states' in kwargs else args[0]
    weight        = kwargs.pop('weight')        if 'weight'        in kwargs else args[1] if len(args) > 1 else None
    if weight is None:
        raise TypeError("run() missing required argument: 'weight'")
    if kwargs:
        raise TypeError(f"run() got unexpected keyword arguments {list(kwargs.keys())}")

    # -------------------------------------------------------------------------#
    # Shape / dtype checks
    # -------------------------------------------------------------------------#
    if hidden_states.ndim != 2:
        raise ValueError("hidden_states must be 2-D [batch, hidden]")
    batch, hidden = hidden_states.shape
    if hidden != HIDDEN_SIZE:
        raise ValueError(f"hidden dimension must be {HIDDEN_SIZE}")
    if weight.shape != (HIDDEN_SIZE,):
        raise ValueError(f"weight shape must be ({HIDDEN_SIZE},)")

    # -------------------------------------------------------------------------#
    # Device handling
    # -------------------------------------------------------------------------#
    if not torch.cuda.is_available():
        if hidden_states.is_cuda or weight.is_cuda:
            raise RuntimeError("CUDA tensors provided but CUDA is not available")
        # CPU fallback (reference implementation)
        x = hidden_states.to(torch.float32)
        inv_rms = torch.rsqrt(x.pow(2).mean(dim=-1, keepdim=True) + EPS)
        y = (x * inv_rms) * weight.to(torch.float32)
        return y.to(hidden_states.dtype)

    orig_device = hidden_states.device
    x_gpu = hidden_states if hidden_states.is_cuda else hidden_states.cuda()
    w_gpu = weight        if weight.is_cuda        else weight.cuda()

    # Ensure contiguous layout for predictable strides
    x_gpu = x_gpu.contiguous()
    w_gpu = w_gpu.contiguous()

    # Allocate output
    o_gpu = torch.empty_like(x_gpu)

    # -------------------------------------------------------------------------#
    # Kernel launch
    # -------------------------------------------------------------------------#
    grid = (batch,)
    _rmsnorm_kernel[grid](
        x_gpu, w_gpu, o_gpu,
        x_gpu.stride(0), o_gpu.stride(0),
        EPS, HIDDEN_SIZE,
        num_warps=4
    )

    # -------------------------------------------------------------------------#
    # Move back to original device if necessary
    # -------------------------------------------------------------------------#
    if orig_device.type == 'cpu':
        return o_gpu.cpu()
    return o_gpu


# -----------------------------------------------------------------------------#
# This file exposes a single callable `run` for external use
# -----------------------------------------------------------------------------#
__all__ = ["run"]
scrolls · 143 lines total

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

Best evidence level for this revision: reported

JSON