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
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 = 4
num_warps=4, num_stages=4stages = 4
num_warps=4, num_stages=4Kernel 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_devscrolls · 116 lines total
Source code from FlashInfer-Bench (flashinfer-ai/flashinfer-trace) · Apache-2.0
Best evidence level for this revision: reported
JSON