gpt-o3 / tritone5fd98
gpt-o3_triton_e5fd98 · gpt-o3 · triton · Apache-2.0
Use it
Vendorable · source mirrored · Apache-2.0View source →
No package. Vendor the mirrored source: 108 lines, Apache-2.0, pinned at da91508.
main.py
curl "https://kernelindex.com/api/v1/implementations/flashinfer-gpt-o3-triton-e5fd98?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:a23031c5b9cec36a07ecf730e0c60e814d3e5045f62941b1fa704e4ed9d5c775
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=8Kernel source
main.py108 lines
import math
import torch
import triton
import triton.language as tl
# ----------------------------------------------------------------------------- #
# Constants
# ----------------------------------------------------------------------------- #
_HIDDEN_SIZE: int = 7168 # fixed hidden dimension
_BLOCK_SIZE: int = 256 # elements processed per kernel instance
_EPS_DEFAULT: float = 1e-6 # numerical stability
# ----------------------------------------------------------------------------- #
# Triton kernel
# ----------------------------------------------------------------------------- #
@triton.jit
def _fused_add_rmsnorm_h7168_kernel(
hidden_ptr, # *bf16 [batch, 7168]
residual_ptr, # *bf16 [batch, 7168]
weight_ptr, # *bf16 [7168]
out_ptr, # *bf16 [batch, 7168]
eps, # float32
hidden_size: tl.constexpr, # 7168 (compile-time constant)
BLOCK_SIZE: tl.constexpr # 256 (compile-time constant)
):
pid = tl.program_id(axis=0) # row / batch index
row_start = pid * hidden_size # base pointer for the row
offs = tl.arange(0, BLOCK_SIZE) # element indices inside the block
# --------------------------------------------------------------------- #
# Pass 1 : compute mean of squares
# --------------------------------------------------------------------- #
sum_sq = tl.zeros((), tl.float32)
for offset in tl.static_range(0, hidden_size, BLOCK_SIZE):
idx = offset + offs
mask = idx < hidden_size
h = tl.load(hidden_ptr + row_start + idx, mask=mask, other=0).to(tl.float32)
r = tl.load(residual_ptr + row_start + idx, mask=mask, other=0).to(tl.float32)
x = h + r
sum_sq += tl.sum(x * x, axis=0)
mean_sq = sum_sq / hidden_size
inv_rms = tl.rsqrt(mean_sq + eps)
# --------------------------------------------------------------------- #
# Pass 2 : normalize, scale and store
# --------------------------------------------------------------------- #
for offset in tl.static_range(0, hidden_size, BLOCK_SIZE):
idx = offset + offs
mask = idx < hidden_size
h = tl.load(hidden_ptr + row_start + idx, mask=mask, other=0).to(tl.float32)
r = tl.load(residual_ptr + row_start + idx, mask=mask, other=0).to(tl.float32)
w = tl.load(weight_ptr + idx, mask=mask, other=0).to(tl.float32)
y = (h + r) * inv_rms * w
tl.store(out_ptr + row_start + idx, y.to(tl.bfloat16), mask=mask)
# ----------------------------------------------------------------------------- #
# Wrapper
# ----------------------------------------------------------------------------- #
def run(hidden_states: torch.Tensor,
residual: torch.Tensor,
weight: torch.Tensor,
eps: float = _EPS_DEFAULT) -> torch.Tensor:
"""
Fused Add + RMSNorm for hidden size 7168 (BF16) using Triton on B200 GPUs.
"""
# ---------------------------- Sanity checks --------------------------- #
if hidden_states.shape != residual.shape:
raise ValueError("hidden_states and residual must have identical shapes")
if hidden_states.shape[-1] != _HIDDEN_SIZE:
raise ValueError(f"hidden_size must be {_HIDDEN_SIZE}, "
f"got {hidden_states.shape[-1]}")
if not torch.cuda.is_available():
raise RuntimeError("CUDA is not available — cannot execute Triton kernel")
# ------------------------- Device management -------------------------- #
orig_device = hidden_states.device # remember caller's device
device = torch.device("cuda") # execute on default CUDA
hidden_gpu = hidden_states.to(device, dtype=torch.bfloat16, non_blocking=True)
residual_gpu = residual.to(device, dtype=torch.bfloat16, non_blocking=True)
weight_gpu = weight.to(device, dtype=torch.bfloat16, non_blocking=True)
batch_size = hidden_gpu.shape[0]
output_gpu = torch.empty_like(hidden_gpu, device=device, dtype=torch.bfloat16)
# ------------------------- Launch Triton kernel ----------------------- #
grid = (batch_size,)
_fused_add_rmsnorm_h7168_kernel[grid](
hidden_gpu,
residual_gpu,
weight_gpu,
output_gpu,
eps,
hidden_size=_HIDDEN_SIZE,
BLOCK_SIZE=_BLOCK_SIZE,
num_warps=8
)
# ----------------------------- Return --------------------------------- #
if orig_device.type == "cuda":
return output_gpu.to(orig_device, non_blocking=True)
return output_gpu.cpu()scrolls · 108 lines total
Source code from FlashInfer-Bench (flashinfer-ai/flashinfer-trace) · Apache-2.0
Best evidence level for this revision: reported
JSON