gpt-o3 / triton35b90e
gpt-o3_triton_35b90e · gpt-o3 · triton · Apache-2.0
Kernel source · 143 lines ↓holds 1 record
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
Show all 14 measurements ›Showing all 14 measurements ⌄
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 = 4
num_warps=4Kernel 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