Skip to content
KernelIndex
Search⌘K

gpt-o3 / tritonad56c1

gpt-o3_triton_ad56c1 · gpt-o3 · triton · Apache-2.0

Use it

Vendorable · source mirrored · Apache-2.0View source →

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

main.py
curl "https://kernelindex.com/api/v1/implementations/flashinfer-gpt-o3-triton-ad56c1?include=source"
interfacetriton
revisionda915083d4c7
symbolrun
pathmain.py
Compatibility
measured onNVIDIA B200
declared hardwareNVIDIA B200
architecturessm_100
dtypesbf16, fp32, int32

Benchmark evidence

38 measurements across 1 GPU, fastest first.

Operation / workload
Hardware
Latency
Rank
Observed
NVIDIA B200
144.3µs
#4 of 5
2025-10-16
NVIDIA B200
146.7µs
#4 of 5
2025-10-16
NVIDIA B200
149.7µs
#4 of 5
2025-10-16
NVIDIA B200
151.8µs
#4 of 5
2025-10-16
NVIDIA B200
153.4µs
#4 of 5
2025-10-16
NVIDIA B200
157.3µs
#4 of 5
2025-10-16
NVIDIA B200
159.3µs
#4 of 5
2025-10-16
NVIDIA B200
160.8µs
#4 of 5
2025-10-16
NVIDIA B200
176.9µs
#4 of 5
2025-10-16
NVIDIA B200
181.4µs
#4 of 5
2025-10-16
Show all 38 measurements ›
NVIDIA B200
187.1µs
#4 of 5
2025-10-16
NVIDIA B200
200.8µs
#4 of 5
2025-10-16
NVIDIA B200
258.3µs
#4 of 5
2025-10-16
NVIDIA B200
273.7µs
#4 of 5
2025-10-16
NVIDIA B200
309.2µs
#4 of 5
2025-10-16
NVIDIA B200
406.3µs
#4 of 5
2025-10-16
NVIDIA B200
455.9µs
#4 of 5
2025-10-16
NVIDIA B200
482.1µs
#4 of 5
2025-10-16
NVIDIA B200
508.0µs
#4 of 5
2025-10-16
NVIDIA B200
646.8µs
#4 of 5
2025-10-16
NVIDIA B200
683.2µs
#4 of 5
2025-10-16
NVIDIA B200
707.2µs
#4 of 5
2025-10-16
NVIDIA B200
1.19ms
#4 of 5
2025-10-16
NVIDIA B200
1.94ms
#4 of 5
2025-10-16
NVIDIA B200
2.33ms
#4 of 5
2025-10-16
NVIDIA B200
4.73ms
#3 of 5
2025-10-16
NVIDIA B200
6.95ms
#4 of 5
2025-10-16
NVIDIA B200
8.83ms
#4 of 5
2025-10-16
NVIDIA B200
13.3ms
#4 of 5
2025-10-16
NVIDIA B200
14.6ms
#4 of 5
2025-10-16
NVIDIA B200
34.7ms
#4 of 5
2025-10-16
NVIDIA B200
85.3ms
#4 of 5
2025-10-16
NVIDIA B200
89.9ms
#4 of 5
2025-10-16
NVIDIA B200
148.6ms
#4 of 5
2025-10-16
NVIDIA B200
329.5ms
#4 of 5
2025-10-16
NVIDIA B200
462.1ms
#4 of 5
2025-10-16
NVIDIA B200
713.5ms
#4 of 4
2025-10-16
NVIDIA B200
3.23s
#4 of 4
2025-10-16

Reported · How evidence levels are derived →

Source and license

sourcehttps://huggingface.co/datasets/flashinfer-ai/flashinfer-trace
commitda915083d4c7c5e61aa3005e3d17ae488e0fc71c
revision digestsha256:ad057faa231033181e08f35d96fbd208a61042e2dd0ef32835e77da5486c5ba8
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 = 4num_warps=4, num_stages=2,
online-softmaxm_new = tl.maximum(m_prev, m_blk)
stages = 2num_warps=4, num_stages=2,
tile-k = 32BLOCK_K: tl.constexpr = 32,

Kernel source

main.py206 lines
import math
import torch
import triton
import triton.language as tl


# ============================================================================ #
#                              TRITON KERNEL                                   #
# ============================================================================ #
@triton.jit
def _mla_paged_prefill_kernel(
    q_nope_ptr,          # *bf16 [TOTAL_Q, 16, 512]
    q_pe_ptr,            # *bf16 [TOTAL_Q, 16,  64]
    kc_ptr,              # *bf16 [KV_LEN,        512]  (sequence-contiguous)
    kp_ptr,              # *bf16 [KV_LEN,         64]
    out_ptr,             # *bf16 [TOTAL_Q, 16, 512]
    lse_ptr,             # *fp32 [TOTAL_Q, 16]
    kv_len,              #   i32 – #tokens in this sequence’s KV buffer
    prefix_len,          #   i32 – kv_len - q_len
    sm_scale,            #  fp32 – soft-max scale
    q_global_offset,     #   i32 – start row of this sequence in Q tensors
    BLOCK_K:  tl.constexpr = 32,
    HEAD_C:   tl.constexpr = 512,
    HEAD_P:   tl.constexpr = 64,
):
    """
    One program instance = (one query token, one head).

    Grid  = (q_len, 16)
    Implements streaming softmax with causal masking and fused output.
    """

    # --------------------------------------------------------------------- #
    #                               INDICES                                 #
    # --------------------------------------------------------------------- #
    pid_q = tl.program_id(0)          # query index inside the sequence
    pid_h = tl.program_id(1)          # head (0‥15)

    q_row = q_global_offset + pid_q   # absolute query row inside Q tensors

    # --------------------------------------------------------------------- #
    #                           LOAD QUERY VECTORS                          #
    # --------------------------------------------------------------------- #
    qn_off = (q_row * 16 + pid_h) * HEAD_C + tl.arange(0, HEAD_C)
    qp_off = (q_row * 16 + pid_h) * HEAD_P + tl.arange(0, HEAD_P)

    qn = tl.load(q_nope_ptr + qn_off).to(tl.float32)   # [512]
    qp = tl.load(q_pe_ptr   + qp_off).to(tl.float32)   # [ 64]

    # --------------------------------------------------------------------- #
    #                      STREAMING SOFTMAX ACCUMULATORS                   #
    # --------------------------------------------------------------------- #
    m_prev  = tl.full((), -float("inf"), tl.float32)        # running max
    l_prev  = tl.zeros((), tl.float32)                      # running sum(exp)
    acc_out = tl.zeros((HEAD_C,), tl.float32)               # running numerator

    query_abs_pos = prefix_len + pid_q                      # absolute position
    num_iters     = (kv_len + BLOCK_K - 1) // BLOCK_K

    INV_LN2 = 1.4426950408889634                            # 1 / ln(2)

    iter_idx = 0
    while iter_idx < num_iters:
        k_start  = iter_idx * BLOCK_K
        tok_offs = k_start + tl.arange(0, BLOCK_K)          # [B]
        valid_m  = tok_offs < kv_len                        # [B]
        causal_m = tok_offs > query_abs_pos                 # [B]
        keep_m   = valid_m & ~causal_m                      # [B]

        # ------------------------ LOAD KC / KP BLOCK --------------------- #
        kc_ptrs = kc_ptr + tok_offs[:, None] * HEAD_C + tl.arange(0, HEAD_C)[None, :]
        kp_ptrs = kp_ptr + tok_offs[:, None] * HEAD_P + tl.arange(0, HEAD_P)[None, :]

        kc_blk = tl.load(kc_ptrs, mask=valid_m[:, None], other=0).to(tl.float32)  # [B,512]
        kp_blk = tl.load(kp_ptrs, mask=valid_m[:, None], other=0).to(tl.float32)  # [B, 64]

        # ---------------------------- DOTS ------------------------------- #
        dotkc  = tl.sum(kc_blk * qn[None, :], 1)     # [B]
        dotkp  = tl.sum(kp_blk * qp[None, :], 1)     # [B]
        logits = (dotkc + dotkp) * sm_scale          # [B]

        neg_inf = -float("inf")
        logits  = tl.where(keep_m, logits, neg_inf)

        # ----------------------- STABLE SOFTMAX -------------------------- #
        m_blk = tl.max(logits, 0)
        m_new = tl.maximum(m_prev, m_blk)

        exp_logits = tl.exp(logits - m_new)
        exp_logits = tl.where(keep_m, exp_logits, 0.0)

        alpha_prev = tl.exp(m_prev - m_new)
        l_prev     = l_prev * alpha_prev + tl.sum(exp_logits, 0)
        acc_out    = acc_out * alpha_prev + tl.sum(exp_logits[:, None] * kc_blk, 0)

        m_prev   = m_new
        iter_idx += 1

    # --------------------------- WRITE-BACK ------------------------------ #
    out_vec  = acc_out / l_prev
    out_ptrs = out_ptr + (q_row * 16 + pid_h) * HEAD_C + tl.arange(0, HEAD_C)
    tl.store(out_ptrs, out_vec.to(tl.bfloat16))

    lse_val  = (tl.log(l_prev) + m_prev) * INV_LN2          # base-2 log-sum-exp
    tl.store(lse_ptr + q_row * 16 + pid_h, lse_val)


# ============================================================================ #
#                              PYTHON ENTRY                                    #
# ============================================================================ #
def run(
    q_nope,          # bf16 [total_q , 16, 512]
    q_pe,            # bf16 [total_q , 16,  64]
    ckv_cache,       # bf16 [num_pages, 1, 512]
    kpe_cache,       # bf16 [num_pages, 1,  64]
    qo_indptr,       # int32 [len_indptr]
    kv_indptr,       # int32 [len_indptr]
    kv_indices,      # int32 [num_kv_indices]
    sm_scale=None,   # optional float
):
    """
    Optimized paged-KV prefill for B200 GPUs.
    Handles device transfers transparently.
    """
    if not torch.cuda.is_available():
        raise RuntimeError("CUDA device is required but not available.")

    NUM_HEADS  = 16
    HEAD_C     = 512
    HEAD_P     = 64
    PAGE_SIZE  = 1
    BLOCK_K    = 32   # must stay in sync with kernel default

    # --------------------------- SHAPE CHECKS --------------------------- #
    assert q_nope.shape[1:] == (NUM_HEADS, HEAD_C)
    assert q_pe.shape[1:]   == (NUM_HEADS, HEAD_P)
    assert ckv_cache.shape[1] == PAGE_SIZE
    assert kpe_cache.shape[1] == PAGE_SIZE
    assert q_nope.shape[0] == qo_indptr[-1].item()
    assert kv_indices.shape[0] == kv_indptr[-1].item()

    # ---------------------------- DEVICE I/O --------------------------- #
    def _to_cuda(t: torch.Tensor):
        return t.cuda(non_blocking=True) if t.device.type != "cuda" else t

    def _back(t: torch.Tensor, ref: torch.Tensor):
        return t.cpu() if ref.device.type != "cuda" else t

    q_nope_c   = _to_cuda(q_nope)
    q_pe_c     = _to_cuda(q_pe)
    kc_all     = _to_cuda(ckv_cache).squeeze(1).contiguous()   # [pages,512]
    kp_all     = _to_cuda(kpe_cache).squeeze(1).contiguous()   # [pages, 64]
    qo_ind_c   = _to_cuda(qo_indptr)
    kv_ind_c   = _to_cuda(kv_indptr)
    kv_idx_c   = _to_cuda(kv_indices)

    total_q = q_nope_c.shape[0]
    batch   = qo_ind_c.shape[0] - 1

    # ------------------------- OUTPUT BUFFERS --------------------------- #
    out_c = torch.empty_like(q_nope_c)
    lse_c = torch.full(
        (total_q, NUM_HEADS), -float("inf"), dtype=torch.float32, device=q_nope_c.device
    )

    # --------------------------- SM SCALE ------------------------------ #
    if sm_scale is None:
        sm_scale = 1.0 / math.sqrt(HEAD_C)
    sm_scale = float(sm_scale)

    # ------------------------- SEQUENCE LOOP --------------------------- #
    for b in range(batch):
        q_beg, q_end = int(qo_ind_c[b].item()), int(qo_ind_c[b + 1].item())
        if q_beg >= q_end:
            continue

        p_beg, p_end = int(kv_ind_c[b].item()), int(kv_ind_c[b + 1].item())
        if p_beg >= p_end:
            continue

        kv_pages = kv_idx_c[p_beg:p_end].long()
        kv_len   = kv_pages.numel()
        q_len    = q_end - q_beg
        prefix   = kv_len - q_len
        if prefix < 0:
            raise RuntimeError("KV length must be ≥ query length (causal)")

        # Gather contiguous KC / KP for this sequence
        kc_seq = kc_all.index_select(0, kv_pages).contiguous()
        kp_seq = kp_all.index_select(0, kv_pages).contiguous()

        grid = (q_len, NUM_HEADS)  # (pid_q, pid_h)

        _mla_paged_prefill_kernel[grid](
            q_nope_c, q_pe_c,
            kc_seq, kp_seq,
            out_c, lse_c,
            kv_len, prefix,
            sm_scale, q_beg,
            num_warps=4, num_stages=2,
            BLOCK_K=BLOCK_K,
            HEAD_C=HEAD_C,
            HEAD_P=HEAD_P,
        )

    return _back(out_c, q_nope), _back(lse_c, q_nope)
scrolls · 206 lines total

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

Best evidence level for this revision: reported

JSON