gpt-o3 / triton2e18c2
gpt-o3_triton_2e18c2 · gpt-o3 · triton · Apache-2.0
Kernel source · 87 lines ↓holds 4 records
Use it
Vendorable · source mirrored · Apache-2.0View source →
No package. Vendor the mirrored source: 87 lines, Apache-2.0, pinned at da91508.
main.py
curl "https://kernelindex.com/api/v1/implementations/flashinfer-gpt-o3-triton-2e18c2?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:6268b10b463aba9d7a16107cc449ff39dc585978703c3703b7f3e773ea376b09
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,stages = 4
num_stages=4,Kernel source
main.py87 lines
import torch
import triton
import triton.language as tl
@triton.jit
def _fused_add_rmsnorm_kernel(
hidden_ptr, # *bf16 [batch_size, 2048]
residual_ptr, # *bf16 [batch_size, 2048]
weight_ptr, # *bf16 [2048]
output_ptr, # *bf16 [batch_size, 2048]
batch_size, # int
eps, # float32
BLOCK_SIZE: tl.constexpr, # 2048
):
row = tl.program_id(0)
if row >= batch_size:
return
offs = tl.arange(0, BLOCK_SIZE)
hidden_ptrs = hidden_ptr + row * BLOCK_SIZE + offs
residual_ptrs = residual_ptr + row * BLOCK_SIZE + offs
weight_ptrs = weight_ptr + offs
out_ptrs = output_ptr + row * BLOCK_SIZE + offs
hidden = tl.load(hidden_ptrs).to(tl.float32)
residual = tl.load(residual_ptrs).to(tl.float32)
weight = tl.load(weight_ptrs).to(tl.float32)
x = hidden + residual
sq = x * x
ssq = tl.sum(sq, axis=0)
inv_rms = tl.math.rsqrt(ssq / BLOCK_SIZE + eps)
y = x * inv_rms * weight
y = y.to(tl.bfloat16)
tl.store(out_ptrs, y)
def run(hidden_states, residual, weight, eps=1e-6):
# Basic validation
if hidden_states.ndim != 2 or residual.ndim != 2:
raise ValueError("hidden_states and residual must be 2-D tensors")
if hidden_states.shape != residual.shape:
raise ValueError("hidden_states and residual must have identical shapes")
if hidden_states.shape[1] != 2048 or weight.numel() != 2048:
raise ValueError("hidden_size must be 2048 for all tensors")
if hidden_states.dtype != torch.bfloat16 or residual.dtype != torch.bfloat16 or weight.dtype != torch.bfloat16:
raise ValueError("All tensors must be of dtype torch.bfloat16")
if hidden_states.device != residual.device or hidden_states.device.type != weight.device.type:
raise ValueError("All input tensors must reside on the same device")
batch_size = hidden_states.shape[0]
src_device = hidden_states.device
# Ensure data is on CUDA
if src_device.type != "cuda":
if not torch.cuda.is_available():
raise RuntimeError("CUDA is not available but is required for Triton kernels")
device = torch.device("cuda")
hidden_cuda = hidden_states.to(device)
residual_cuda = residual.to(device)
weight_cuda = weight.to(device)
else:
device = src_device
hidden_cuda = hidden_states.contiguous()
residual_cuda = residual.contiguous()
weight_cuda = weight.contiguous()
output_cuda = torch.empty_like(hidden_cuda)
grid = (batch_size,)
_fused_add_rmsnorm_kernel[grid](
hidden_cuda,
residual_cuda,
weight_cuda,
output_cuda,
batch_size,
eps,
BLOCK_SIZE=2048,
num_warps=8,
num_stages=4,
)
return output_cuda if src_device.type == "cuda" else output_cuda.to(src_device)scrolls · 87 lines total
Source code from FlashInfer-Bench (flashinfer-ai/flashinfer-trace) · Apache-2.0
Best evidence level for this revision: reported
JSON