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
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 = 8
NUM_WARPS = 8 # good occupancy for B200stages = 1
num_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