gpt-o3 / triton19c647
gpt-o3_triton_19c647 · gpt-o3 · triton · Apache-2.0
Kernel source · 112 lines ↓holds 3 records
Use it
Vendorable · source mirrored · Apache-2.0View source →
No package. Vendor the mirrored source: 112 lines, Apache-2.0, pinned at da91508.
main.py
curl "https://kernelindex.com/api/v1/implementations/flashinfer-gpt-o3-triton-19c647?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:398e112657280d9281ed964554b3ca2a1cafdca6dcd1e5ceca14979c2bb5d978
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.py112 lines
import math
import torch
import triton
import triton.language as tl
# ----------------------------------------------------------------------
# Triton kernel
# ----------------------------------------------------------------------
@triton.jit
def _rmsnorm_kernel(
x_ptr, # pointer to hidden_states
w_ptr, # pointer to weight
o_ptr, # pointer to output
stride_bs, # hidden_size stride between rows
eps, # epsilon for numerical stability
BLOCK_SIZE: tl.constexpr # hidden dimension (constant = 512)
):
batch_id = tl.program_id(0)
offs = tl.arange(0, BLOCK_SIZE)
# Pointers for this row
x_row_ptrs = x_ptr + batch_id * stride_bs + offs
w_ptrs = w_ptr + offs
o_row_ptrs = o_ptr + batch_id * stride_bs + offs
# Load inputs as bf16 then upcast to fp32 for compute
x = tl.load(x_row_ptrs).to(tl.float32)
w = tl.load(w_ptrs).to(tl.float32)
# Compute inverse RMS
mean_square = tl.sum(x * x, axis=0) / BLOCK_SIZE
inv_rms = tl.math.rsqrt(mean_square + eps)
# Normalize and scale
y = x * inv_rms * w
# Store result (down-cast to bf16)
tl.store(o_row_ptrs, y.to(tl.bfloat16))
# ----------------------------------------------------------------------
# Python wrapper
# ----------------------------------------------------------------------
@torch.no_grad()
def run(hidden_states: torch.Tensor,
weight: torch.Tensor,
eps: float = 1e-6):
"""
RMSNorm (hidden_size = 512) implemented with a Triton kernel.
Args:
hidden_states: [batch_size, 512] BF16 tensor
weight: [512] BF16 tensor
eps: float, epsilon used for numerical stability
Returns:
output: [batch_size, 512] BF16 tensor (same device as input)
"""
# ------------------------------------------------------------------
# Sanity checks
# ------------------------------------------------------------------
if hidden_states.dim() != 2:
raise ValueError("hidden_states must be 2-D [batch_size, 512]")
if weight.dim() != 1 or weight.numel() != 512:
raise ValueError("weight must be 1-D with 512 elements")
if hidden_states.shape[1] != 512:
raise ValueError("hidden_size is fixed to 512")
if not torch.cuda.is_available():
raise RuntimeError("CUDA device is required to run Triton kernel")
# ------------------------------------------------------------------
# Device management
# ------------------------------------------------------------------
target_device = hidden_states.device
if target_device.type == "cpu":
target_device = torch.device("cuda")
hidden_states_gpu = hidden_states.to(device=target_device, dtype=torch.bfloat16, copy=False)
weight_gpu = weight.to(device=target_device, dtype=torch.bfloat16, copy=False)
# Output allocation
output_gpu = torch.empty_like(hidden_states_gpu)
# ------------------------------------------------------------------
# Kernel launch configuration
# ------------------------------------------------------------------
batch_size = hidden_states_gpu.shape[0]
BLOCK_SIZE = 512
grid = (batch_size,)
_rmsnorm_kernel[grid](
hidden_states_gpu,
weight_gpu,
output_gpu,
hidden_states_gpu.stride(0),
eps,
BLOCK_SIZE=BLOCK_SIZE,
num_warps=8,
num_stages=4,
)
# ------------------------------------------------------------------
# Return result on original device
# ------------------------------------------------------------------
if hidden_states.device.type == "cpu":
return output_gpu.cpu()
return output_gpu
# ----------------------------------------------------------------------
# Export for external use
# ----------------------------------------------------------------------
__all__ = ["run"]scrolls · 112 lines total
Source code from FlashInfer-Bench (flashinfer-ai/flashinfer-trace) · Apache-2.0
Best evidence level for this revision: reported
JSON