Skip to content
KernelIndex
Search⌘K

gpt-o3 / tritonc1e819

gpt-o3_triton_c1e819 · gpt-o3 · triton · Apache-2.0

Use it

Vendorable · source mirrored · Apache-2.0View source →

No package. Vendor the mirrored source: 128 lines, Apache-2.0, pinned at da91508.

main.py
curl "https://kernelindex.com/api/v1/implementations/flashinfer-gpt-o3-triton-c1e819?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
Fused add RMSNorm h4096bf16 · [4096] · batch_size=34
NVIDIA B200
10.2µs
#4 of 8
2025-10-16
Fused add RMSNorm h4096bf16 · [4096] · batch_size=64
NVIDIA B200
10.2µs
#4 of 8
2025-10-16
Fused add RMSNorm h4096bf16 · [4096] · batch_size=16
NVIDIA B200
10.2µs
#4 of 8
2025-10-16
Fused add RMSNorm h4096bf16 · [4096] · batch_size=15
NVIDIA B200
10.2µs
#4 of 8
2025-10-16
Fused add RMSNorm h4096bf16 · [4096] · batch_size=63
NVIDIA B200
10.2µs
#4 of 8
2025-10-16
Fused add RMSNorm h4096bf16 · [4096] · batch_size=7
NVIDIA B200
10.3µs
#4 of 8
2025-10-16
Fused add RMSNorm h4096bf16 · [4096] · batch_size=1
NVIDIA B200
10.3µs
#3 of 8
2025-10-16
Fused add RMSNorm h4096bf16 · [4096] · batch_size=79
NVIDIA B200
10.3µs
#4 of 8
2025-10-16
Fused add RMSNorm h4096bf16 · [4096] · batch_size=170
NVIDIA B200
10.4µs
#4 of 8
2025-10-16
Fused add RMSNorm h4096bf16 · [4096] · batch_size=8804
NVIDIA B200
53.4µs
#2 of 8
2025-10-16
Show all 14 measurements ›
Fused add RMSNorm h4096bf16 · [4096] · batch_size=10827
NVIDIA B200
63.5µs
#3 of 8
2025-10-16
Fused add RMSNorm h4096bf16 · [4096] · batch_size=11832
NVIDIA B200
67.8µs
#2 of 8
2025-10-16
Fused add RMSNorm h4096bf16 · [4096] · batch_size=14418
NVIDIA B200
79.9µs
#3 of 8
2025-10-16
Fused add RMSNorm h4096bf16 · [4096] · batch_size=14509
NVIDIA B200
80.0µ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:e25739502af57a6d7db7a9290eb16603f9057261e4bf11647f4e17fa984bfc5e
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.py128 lines
import torch
import triton
import triton.language as tl

# ----------------------------------------------------------------------
# Problem-wide constants (compile-time)
# ----------------------------------------------------------------------
HIDDEN_SIZE: int = 4096          # fixed hidden dimension
BLOCK_SIZE:  int = 1024          # 4 chunks per row
EPSILON:     float = 1e-5        # epsilon for RMSNorm


# ----------------------------------------------------------------------
# Triton kernel
# ----------------------------------------------------------------------
@triton.jit
def _fused_add_rmsnorm_h4096(
    hidden_ptr,        # bf16 [B, H]
    residual_ptr,      # bf16 [B, H]
    weight_ptr,        # bf16 [H]
    output_ptr,        # bf16 [B, H]
    stride_hidden_bs,  # = HIDDEN_SIZE
    stride_res_bs,     # = HIDDEN_SIZE
    stride_out_bs,     # = HIDDEN_SIZE
    BLOCK: tl.constexpr = BLOCK_SIZE,
    H:      tl.constexpr = HIDDEN_SIZE,
    EPS:    tl.constexpr = EPSILON,
):
    """
    One program instance handles one row (batch element).
    It iterates over the hidden dimension in BLOCK-wide chunks.
    """

    # Program/id along batch dimension
    pid = tl.program_id(axis=0)

    # Offsets for a block of columns [0, BLOCK)
    offs = tl.arange(0, BLOCK)

    # Base pointers for this row
    hidden_row   = hidden_ptr   + pid * stride_hidden_bs
    residual_row = residual_ptr + pid * stride_res_bs
    output_row   = output_ptr   + pid * stride_out_bs

    # ------------------------------------------------------------------
    # Pass 1: compute sum of squares
    # ------------------------------------------------------------------
    ssq = tl.zeros((), dtype=tl.float32)

    for col in tl.static_range(0, H, BLOCK):
        idx  = col + offs
        x_h  = tl.load(hidden_row   + idx).to(tl.float32)
        x_r  = tl.load(residual_row + idx).to(tl.float32)
        x    = x_h + x_r
        ssq += tl.sum(x * x, axis=0)

    inv_rms = tl.rsqrt(ssq / H + EPS)

    # ------------------------------------------------------------------
    # Pass 2: normalize, scale and store
    # ------------------------------------------------------------------
    for col in tl.static_range(0, H, BLOCK):
        idx  = col + offs
        h    = tl.load(hidden_row   + idx).to(tl.float32)
        r    = tl.load(residual_row + idx).to(tl.float32)
        w    = tl.load(weight_ptr    + idx).to(tl.float32)

        y = (h + r) * inv_rms * w
        tl.store(output_row + idx, y.to(tl.bfloat16))


# ----------------------------------------------------------------------
# Python wrapper
# ----------------------------------------------------------------------
def run(hidden_states: torch.Tensor,
        residual:       torch.Tensor,
        weight:         torch.Tensor,
        **kwargs) -> torch.Tensor:
    """
    Fused Add + RMSNorm (hidden_size = 4096, bf16) implemented in Triton.
    """

    # ------------------ 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}")
    if weight.numel() != HIDDEN_SIZE:
        raise ValueError(f"weight must have {HIDDEN_SIZE} elements")

    if not torch.cuda.is_available():
        raise RuntimeError("CUDA device is required to run Triton kernels")

    # ------------------ Device handling ----------------
    src_device = hidden_states.device
    device = src_device if src_device.type == "cuda" else torch.device("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)

    # Ensure contiguous layout
    hidden_gpu   = hidden_gpu.contiguous()
    residual_gpu = residual_gpu.contiguous()
    weight_gpu   = weight_gpu.contiguous()

    # ------------------ Launch kernel ------------------
    batch_size = hidden_gpu.shape[0]
    output_gpu = torch.empty_like(hidden_gpu)

    grid = (batch_size,)

    _fused_add_rmsnorm_h4096[grid](
        hidden_gpu, residual_gpu, weight_gpu, output_gpu,
        HIDDEN_SIZE, HIDDEN_SIZE, HIDDEN_SIZE,
        num_warps=8,
        num_stages=4,
    )

    # ------------------ Move result back ---------------
    if src_device.type == "cpu":
        return output_gpu.cpu()
    return output_gpu


# ----------------------------------------------------------------------
# End of file
# ----------------------------------------------------------------------
scrolls · 128 lines total

Source code from FlashInfer-Bench (flashinfer-ai/flashinfer-trace) · Apache-2.0

Best evidence level for this revision: reported

JSON