Skip to content
KernelIndex
Search⌘K

gpt-o3 / triton2e18c2

gpt-o3_triton_2e18c2 · gpt-o3 · triton · Apache-2.0

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
Fused add RMSNorm h2048bf16 · [2048] · batch_size=79
NVIDIA B200
6.19µs
#1 of 8
2025-10-16
Fused add RMSNorm h2048bf16 · [2048] · batch_size=64
NVIDIA B200
6.20µs
#1 of 8
2025-10-16
Fused add RMSNorm h2048bf16 · [2048] · batch_size=6
NVIDIA B200
6.21µs
#1 of 8
2025-10-16
Fused add RMSNorm h2048bf16 · [2048] · batch_size=1
NVIDIA B200
6.25µs
#1 of 8
2025-10-16
Fused add RMSNorm h2048bf16 · [2048] · batch_size=34
NVIDIA B200
6.30µs
#3 of 8
2025-10-16
Fused add RMSNorm h2048bf16 · [2048] · batch_size=12383
NVIDIA B200
31.2µs
#2 of 8
2025-10-16
Fused add RMSNorm h2048bf16 · [2048] · batch_size=16254
NVIDIA B200
38.9µs
#2 of 8
2025-10-16

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 = 8num_warps=8,
stages = 4num_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