Skip to content
KernelIndex
Search⌘K

gpt-o3 / tritonb12b97

gpt-o3_triton_b12b97 · gpt-o3 · triton · Apache-2.0

Use it

Vendorable · source mirrored · Apache-2.0View source →

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

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

Benchmark evidence

92 measurements across 1 GPU, fastest first.

Operation / workload
Hardware
Latency
Rank
Observed
GQA paged decode h32 kv4 d128 ps1bf16 · [1, 32, 128] · num_pages=15 · num_kv_indices=14
NVIDIA B200
86.6µs
#6 of 6
2026-03-28
GQA paged decode h32 kv4 d128 ps1bf16 · [1, 32, 128] · num_pages=81 · num_kv_indices=64
NVIDIA B200
87.0µs
#4 of 5
2026-03-28
GQA paged decode h32 kv4 d128 ps1bf16 · [1, 32, 128] · num_pages=9332 · num_kv_indices=87
NVIDIA B200
87.1µs
#2 of 5
2026-03-28
GQA paged decode h32 kv4 d128 ps1bf16 · [1, 32, 128] · num_pages=8 · num_kv_indices=7
NVIDIA B200
87.7µs
#6 of 6
2026-03-28
GQA paged decode h32 kv4 d128 ps1bf16 · [1, 32, 128] · num_pages=17 · num_kv_indices=2
NVIDIA B200
87.8µs
#5 of 5
2026-03-28
GQA paged decode h32 kv4 d128 ps1bf16 · [1, 32, 128] · num_pages=10 · num_kv_indices=9
NVIDIA B200
88.2µs
#5 of 5
2026-03-28
GQA paged decode h32 kv4 d128 ps1bf16 · [1, 32, 128] · num_pages=12 · num_kv_indices=11
NVIDIA B200
88.7µs
#5 of 5
2026-03-28
GQA paged decode h32 kv4 d128 ps1bf16 · [1, 32, 128] · num_pages=191 · num_kv_indices=141
NVIDIA B200
89.3µs
#3 of 6
2026-03-28
GQA paged decode h32 kv4 d128 ps1bf16 · [1, 32, 128] · num_pages=9317 · num_kv_indices=72
NVIDIA B200
89.4µs
#3 of 5
2026-03-28
GQA paged decode h32 kv4 d128 ps1bf16 · [1, 32, 128] · num_pages=57 · num_kv_indices=40
NVIDIA B200
90.3µs
#6 of 6
2026-03-28
Show all 92 measurements ›
GQA paged decode h32 kv4 d128 ps1bf16 · [1, 32, 128] · num_pages=8 · num_kv_indices=7
NVIDIA B200
92.6µs
#7 of 7
2025-10-16
GQA paged decode h32 kv4 d128 ps1bf16 · [1, 32, 128] · num_pages=81 · num_kv_indices=64
NVIDIA B200
93.1µs
#6 of 7
2025-10-16
GQA paged decode h32 kv4 d128 ps1bf16 · [1, 32, 128] · num_pages=17 · num_kv_indices=2
NVIDIA B200
93.6µs
#7 of 7
2025-10-16
GQA paged decode h32 kv4 d128 ps1bf16 · [1, 32, 128] · num_pages=15 · num_kv_indices=14
NVIDIA B200
93.7µs
#7 of 7
2025-10-16
GQA paged decode h32 kv4 d128 ps1bf16 · [1, 32, 128] · num_pages=302 · num_kv_indices=252
NVIDIA B200
94.7µs
#4 of 7
2025-10-16
GQA paged decode h32 kv4 d128 ps1bf16 · [1, 32, 128] · num_pages=57 · num_kv_indices=40
NVIDIA B200
95.5µs
#7 of 7
2025-10-16
GQA paged decode h32 kv4 d128 ps1bf16 · [1, 32, 128] · num_pages=10 · num_kv_indices=9
NVIDIA B200
96.5µs
#7 of 7
2025-10-16
GQA paged decode h32 kv4 d128 ps1bf16 · [1, 32, 128] · num_pages=9317 · num_kv_indices=72
NVIDIA B200
96.6µs
#5 of 7
2025-10-16
GQA paged decode h32 kv4 d128 ps1bf16 · [1, 32, 128] · num_pages=67 · num_kv_indices=50
NVIDIA B200
96.8µs
#7 of 7
2025-10-16
GQA paged decode h32 kv4 d128 ps1bf16 · [1, 32, 128] · num_pages=12 · num_kv_indices=11
NVIDIA B200
97.0µs
#7 of 7
2025-10-16
GQA paged decode h32 kv4 d128 ps1bf16 · [1, 32, 128] · num_pages=191 · num_kv_indices=141
NVIDIA B200
98.3µs
#4 of 7
2025-10-16
GQA paged decode h32 kv4 d128 ps1bf16 · [1, 32, 128] · num_pages=9347 · num_kv_indices=102
NVIDIA B200
99.4µs
#4 of 7
2025-10-16
GQA paged decode h32 kv4 d128 ps1bf16 · [1, 32, 128] · num_pages=412 · num_kv_indices=362
NVIDIA B200
99.7µs
#2 of 5
2026-03-28
GQA paged decode h32 kv4 d128 ps1bf16 · [1, 32, 128] · num_pages=9332 · num_kv_indices=87
NVIDIA B200
104.3µs
#5 of 7
2025-10-16
GQA paged decode h32 kv4 d128 ps1bf16 · [1, 32, 128] · num_pages=412 · num_kv_indices=362
NVIDIA B200
106.3µs
#4 of 7
2025-10-16
GQA paged decode h32 kv4 d128 ps1bf16 · [1, 32, 128] · num_pages=596 · num_kv_indices=546
NVIDIA B200
115.0µs
#3 of 6
2026-03-28
GQA paged decode h32 kv4 d128 ps1bf16 · [1, 32, 128] · num_pages=596 · num_kv_indices=546
NVIDIA B200
116.9µs
#3 of 7
2025-10-16
GQA paged decode h32 kv4 d128 ps1bf16 · [1, 32, 128] · num_pages=486 · num_kv_indices=436
NVIDIA B200
119.8µs
#4 of 7
2025-10-16
GQA paged decode h32 kv4 d128 ps1bf16 · [16, 32, 128] · num_pages=1070 · num_kv_indices=1020
NVIDIA B200
779.6µs
#4 of 4
2026-03-28
GQA paged decode h32 kv4 d128 ps1bf16 · [16, 32, 128] · num_pages=2158 · num_kv_indices=2108
NVIDIA B200
836.2µs
#4 of 4
2026-03-28
GQA paged decode h32 kv4 d128 ps1bf16 · [16, 32, 128] · num_pages=1070 · num_kv_indices=1020
NVIDIA B200
861.8µs
#7 of 7
2025-10-16
GQA paged decode h32 kv4 d128 ps1bf16 · [16, 32, 128] · num_pages=3246 · num_kv_indices=3196
NVIDIA B200
964.8µs
#4 of 4
2026-03-28
GQA paged decode h32 kv4 d128 ps1bf16 · [16, 32, 128] · num_pages=2158 · num_kv_indices=2108
NVIDIA B200
965.6µs
#7 of 7
2025-10-16
GQA paged decode h32 kv4 d128 ps1bf16 · [16, 32, 128] · num_pages=4334 · num_kv_indices=4284
NVIDIA B200
996.1µs
#4 of 4
2026-03-28
GQA paged decode h32 kv4 d128 ps1bf16 · [16, 32, 128] · num_pages=4334 · num_kv_indices=4284
NVIDIA B200
1.06ms
#7 of 7
2025-10-16
GQA paged decode h32 kv4 d128 ps1bf16 · [16, 32, 128] · num_pages=3246 · num_kv_indices=3196
NVIDIA B200
1.07ms
#7 of 7
2025-10-16
GQA paged decode h32 kv4 d128 ps1bf16 · [16, 32, 128] · num_pages=5422 · num_kv_indices=5372
NVIDIA B200
1.15ms
#7 of 7
2025-10-16
GQA paged decode h32 kv4 d128 ps1bf16 · [16, 32, 128] · num_pages=5422 · num_kv_indices=5372
NVIDIA B200
1.15ms
#4 of 4
2026-03-28
GQA paged decode h32 kv4 d128 ps1bf16 · [16, 32, 128] · num_pages=6510 · num_kv_indices=6460
NVIDIA B200
1.18ms
#4 of 4
2026-03-28
GQA paged decode h32 kv4 d128 ps1bf16 · [16, 32, 128] · num_pages=7598 · num_kv_indices=7548
NVIDIA B200
1.20ms
#4 of 4
2026-03-28
GQA paged decode h32 kv4 d128 ps1bf16 · [16, 32, 128] · num_pages=8686 · num_kv_indices=8636
NVIDIA B200
1.24ms
#4 of 4
2026-03-28
GQA paged decode h32 kv4 d128 ps1bf16 · [16, 32, 128] · num_pages=7598 · num_kv_indices=7548
NVIDIA B200
1.28ms
#7 of 7
2025-10-16
GQA paged decode h32 kv4 d128 ps1bf16 · [16, 32, 128] · num_pages=8686 · num_kv_indices=8636
NVIDIA B200
1.30ms
#7 of 7
2025-10-16
GQA paged decode h32 kv4 d128 ps1bf16 · [16, 32, 128] · num_pages=6510 · num_kv_indices=6460
NVIDIA B200
1.43ms
#7 of 7
2025-10-16
GQA paged decode h32 kv4 d128 ps1bf16 · [16, 32, 128] · num_pages=24732 · num_kv_indices=15463
NVIDIA B200
1.66ms
#2 of 4
2026-03-28
GQA paged decode h32 kv4 d128 ps1bf16 · [16, 32, 128] · num_pages=25820 · num_kv_indices=16551
NVIDIA B200
1.74ms
#2 of 4
2026-03-28
GQA paged decode h32 kv4 d128 ps1bf16 · [16, 32, 128] · num_pages=26908 · num_kv_indices=17639
NVIDIA B200
1.77ms
#2 of 4
2026-03-28
GQA paged decode h32 kv4 d128 ps1bf16 · [16, 32, 128] · num_pages=24732 · num_kv_indices=15463
NVIDIA B200
1.78ms
#4 of 7
2025-10-16
GQA paged decode h32 kv4 d128 ps1bf16 · [16, 32, 128] · num_pages=25820 · num_kv_indices=16551
NVIDIA B200
1.79ms
#4 of 7
2025-10-16
GQA paged decode h32 kv4 d128 ps1bf16 · [16, 32, 128] · num_pages=27996 · num_kv_indices=18727
NVIDIA B200
1.88ms
#4 of 7
2025-10-16
GQA paged decode h32 kv4 d128 ps1bf16 · [16, 32, 128] · num_pages=29084 · num_kv_indices=19815
NVIDIA B200
1.91ms
#2 of 4
2026-03-28
GQA paged decode h32 kv4 d128 ps1bf16 · [16, 32, 128] · num_pages=26908 · num_kv_indices=17639
NVIDIA B200
1.91ms
#4 of 7
2025-10-16
GQA paged decode h32 kv4 d128 ps1bf16 · [16, 32, 128] · num_pages=29084 · num_kv_indices=19815
NVIDIA B200
1.96ms
#4 of 7
2025-10-16
GQA paged decode h32 kv4 d128 ps1bf16 · [16, 32, 128] · num_pages=27996 · num_kv_indices=18727
NVIDIA B200
1.99ms
#2 of 4
2026-03-28
GQA paged decode h32 kv4 d128 ps1bf16 · [16, 32, 128] · num_pages=31260 · num_kv_indices=21991
NVIDIA B200
2.06ms
#2 of 4
2026-03-28
GQA paged decode h32 kv4 d128 ps1bf16 · [16, 32, 128] · num_pages=30172 · num_kv_indices=20903
NVIDIA B200
2.09ms
#4 of 7
2025-10-16
GQA paged decode h32 kv4 d128 ps1bf16 · [16, 32, 128] · num_pages=30172 · num_kv_indices=20903
NVIDIA B200
2.10ms
#2 of 4
2026-03-28
GQA paged decode h32 kv4 d128 ps1bf16 · [16, 32, 128] · num_pages=31260 · num_kv_indices=21991
NVIDIA B200
2.11ms
#4 of 7
2025-10-16
GQA paged decode h32 kv4 d128 ps1bf16 · [16, 32, 128] · num_pages=32348 · num_kv_indices=23079
NVIDIA B200
2.14ms
#4 of 7
2025-10-16
GQA paged decode h32 kv4 d128 ps1bf16 · [16, 32, 128] · num_pages=32348 · num_kv_indices=23079
NVIDIA B200
2.22ms
#2 of 4
2026-03-28
GQA paged decode h32 kv4 d128 ps1bf16 · [64, 32, 128] · num_pages=28831 · num_kv_indices=28815
NVIDIA B200
4.45ms
#5 of 5
2026-03-28
GQA paged decode h32 kv4 d128 ps1bf16 · [64, 32, 128] · num_pages=31007 · num_kv_indices=30991
NVIDIA B200
4.57ms
#5 of 5
2026-03-28
GQA paged decode h32 kv4 d128 ps1bf16 · [64, 32, 128] · num_pages=33183 · num_kv_indices=33167
NVIDIA B200
4.74ms
#5 of 5
2026-03-28
GQA paged decode h32 kv4 d128 ps1bf16 · [64, 32, 128] · num_pages=28831 · num_kv_indices=28815
NVIDIA B200
4.80ms
#7 of 7
2025-10-16
GQA paged decode h32 kv4 d128 ps1bf16 · [64, 32, 128] · num_pages=35359 · num_kv_indices=35343
NVIDIA B200
4.94ms
#4 of 4
2026-03-28
GQA paged decode h32 kv4 d128 ps1bf16 · [64, 32, 128] · num_pages=31007 · num_kv_indices=30991
NVIDIA B200
5.03ms
#7 of 7
2025-10-16
GQA paged decode h32 kv4 d128 ps1bf16 · [64, 32, 128] · num_pages=33183 · num_kv_indices=33167
NVIDIA B200
5.06ms
#7 of 7
2025-10-16
GQA paged decode h32 kv4 d128 ps1bf16 · [64, 32, 128] · num_pages=37535 · num_kv_indices=37519
NVIDIA B200
5.14ms
#4 of 4
2026-03-28
GQA paged decode h32 kv4 d128 ps1bf16 · [64, 32, 128] · num_pages=39711 · num_kv_indices=39695
NVIDIA B200
5.21ms
#4 of 4
2026-03-28
GQA paged decode h32 kv4 d128 ps1bf16 · [64, 32, 128] · num_pages=35359 · num_kv_indices=35343
NVIDIA B200
5.25ms
#7 of 7
2025-10-16
GQA paged decode h32 kv4 d128 ps1bf16 · [64, 32, 128] · num_pages=41887 · num_kv_indices=41871
NVIDIA B200
5.37ms
#4 of 4
2026-03-28
GQA paged decode h32 kv4 d128 ps1bf16 · [64, 32, 128] · num_pages=37535 · num_kv_indices=37519
NVIDIA B200
5.42ms
#7 of 7
2025-10-16
GQA paged decode h32 kv4 d128 ps1bf16 · [64, 32, 128] · num_pages=41887 · num_kv_indices=41871
NVIDIA B200
5.55ms
#7 of 7
2025-10-16
GQA paged decode h32 kv4 d128 ps1bf16 · [64, 32, 128] · num_pages=39711 · num_kv_indices=39695
NVIDIA B200
5.67ms
#7 of 7
2025-10-16
GQA paged decode h32 kv4 d128 ps1bf16 · [64, 32, 128] · num_pages=44063 · num_kv_indices=44047
NVIDIA B200
5.67ms
#4 of 4
2026-03-28
GQA paged decode h32 kv4 d128 ps1bf16 · [64, 32, 128] · num_pages=44063 · num_kv_indices=44047
NVIDIA B200
5.67ms
#7 of 7
2025-10-16
GQA paged decode h32 kv4 d128 ps1bf16 · [64, 32, 128] · num_pages=46303 · num_kv_indices=46287
NVIDIA B200
5.78ms
#7 of 7
2025-10-16
GQA paged decode h32 kv4 d128 ps1bf16 · [64, 32, 128] · num_pages=48479 · num_kv_indices=48463
NVIDIA B200
5.96ms
#4 of 4
2026-03-28
GQA paged decode h32 kv4 d128 ps1bf16 · [64, 32, 128] · num_pages=46303 · num_kv_indices=46287
NVIDIA B200
5.96ms
#4 of 4
2026-03-28
GQA paged decode h32 kv4 d128 ps1bf16 · [64, 32, 128] · num_pages=50655 · num_kv_indices=50639
NVIDIA B200
6.04ms
#4 of 4
2026-03-28
GQA paged decode h32 kv4 d128 ps1bf16 · [64, 32, 128] · num_pages=48479 · num_kv_indices=48463
NVIDIA B200
6.08ms
#7 of 7
2025-10-16
GQA paged decode h32 kv4 d128 ps1bf16 · [64, 32, 128] · num_pages=55007 · num_kv_indices=54991
NVIDIA B200
6.18ms
#4 of 4
2026-03-28
GQA paged decode h32 kv4 d128 ps1bf16 · [64, 32, 128] · num_pages=57183 · num_kv_indices=57167
NVIDIA B200
6.25ms
#4 of 4
2026-03-28
GQA paged decode h32 kv4 d128 ps1bf16 · [64, 32, 128] · num_pages=52831 · num_kv_indices=52815
NVIDIA B200
6.33ms
#4 of 4
2026-03-28
GQA paged decode h32 kv4 d128 ps1bf16 · [64, 32, 128] · num_pages=55007 · num_kv_indices=54991
NVIDIA B200
6.38ms
#7 of 7
2025-10-16
GQA paged decode h32 kv4 d128 ps1bf16 · [64, 32, 128] · num_pages=59359 · num_kv_indices=59343
NVIDIA B200
6.40ms
#7 of 7
2025-10-16
GQA paged decode h32 kv4 d128 ps1bf16 · [64, 32, 128] · num_pages=59359 · num_kv_indices=59343
NVIDIA B200
6.42ms
#4 of 4
2026-03-28
GQA paged decode h32 kv4 d128 ps1bf16 · [64, 32, 128] · num_pages=52831 · num_kv_indices=52815
NVIDIA B200
6.46ms
#7 of 7
2025-10-16
GQA paged decode h32 kv4 d128 ps1bf16 · [64, 32, 128] · num_pages=61535 · num_kv_indices=61519
NVIDIA B200
6.73ms
#7 of 7
2025-10-16
GQA paged decode h32 kv4 d128 ps1bf16 · [64, 32, 128] · num_pages=57183 · num_kv_indices=57167
NVIDIA B200
6.74ms
#7 of 7
2025-10-16
GQA paged decode h32 kv4 d128 ps1bf16 · [64, 32, 128] · num_pages=61535 · num_kv_indices=61519
NVIDIA B200
6.84ms
#4 of 4
2026-03-28
GQA paged decode h32 kv4 d128 ps1bf16 · [64, 32, 128] · num_pages=50655 · num_kv_indices=50639
NVIDIA B200
7.54ms
#7 of 7
2025-10-16

Reproduction-ready · How evidence levels are derived →

Source and license

sourcehttps://huggingface.co/datasets/flashinfer-ai/flashinfer-trace
commitda915083d4c7c5e61aa3005e3d17ae488e0fc71c
revision digestsha256:806a80269911f1ea48e10e81d53090bb50b42e8db650c6eebac66fc60534ed90
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.py253 lines
import math
import torch
import triton
import triton.language as tl


# ----------------------------------------------------------------------------- #
#                               Triton Kernel                                   #
# ----------------------------------------------------------------------------- #
@triton.jit
def gqa_paged_decode_kernel(
    q_ptr,                       # *bf16  – [num_qo_heads, head_dim]
    k_cache_ptr,                 # *bf16  – [num_pages, num_kv_heads, head_dim]
    v_cache_ptr,                 # *bf16  – [num_pages, num_kv_heads, head_dim]
    kv_indices_ptr,              # *i32   – [seq_len]
    out_ptr,                     # *bf16  – [num_qo_heads, head_dim]
    lse_ptr,                     # *fp32  – [num_qo_heads]

    sm_scale: tl.float32,        # soft-max scale
    seq_len: tl.constexpr,       # number of tokens for this sequence

    stride_q_head: tl.constexpr,
    stride_k_page: tl.constexpr,
    stride_k_head: tl.constexpr,
    stride_v_page: tl.constexpr,
    stride_v_head: tl.constexpr,
    stride_out_head: tl.constexpr,
    stride_lse_head: tl.constexpr,

    gqa_ratio: tl.constexpr,     # 8 for 32/4
    BLOCK_TOKENS: tl.constexpr,  # 128
    BLOCK_D: tl.constexpr,       # 128
):
    # --------------------------------------------------------------------- #
    # Each program instance processes one query/output (qo) head.
    # --------------------------------------------------------------------- #
    pid   = tl.program_id(axis=0)
    h_qo  = pid                          # 0 … 31
    h_kv  = h_qo // gqa_ratio            # corresponding KV head (0 … 3)

    d_offsets = tl.arange(0, BLOCK_D)

    # --------------------------------------------------------------------- #
    # Load query vector (bf16 → fp32)
    # --------------------------------------------------------------------- #
    q_head_ptr   = q_ptr + h_qo * stride_q_head
    q_vec        = tl.load(q_head_ptr + d_offsets).to(tl.float32)      # [D]

    # Initial values for online softmax reduction
    neg_inf = -1.0e30
    m_val   = tl.full((), neg_inf, dtype=tl.float32)                   # running max

    # --------------------------------------------------------------------- #
    # Pass-1 : find maximum logit
    # --------------------------------------------------------------------- #
    for t0 in range(0, seq_len, BLOCK_TOKENS):
        tok_offsets = tl.arange(0, BLOCK_TOKENS) + t0                  # [T]
        tok_mask    = tok_offsets < seq_len

        page_ids = tl.load(kv_indices_ptr + tok_offsets, mask=tok_mask, other=0)

        k_ptrs = (
            k_cache_ptr
            + page_ids[:, None] * stride_k_page
            + h_kv             * stride_k_head
            + d_offsets[None, :]
        )                                                              # [T, D] pointers

        k_block = tl.load(k_ptrs, mask=tok_mask[:, None], other=0).to(tl.float32)

        logits     = tl.sum(k_block * q_vec[None, :], axis=1) * sm_scale
        logits     = tl.where(tok_mask, logits, neg_inf)
        block_max  = tl.max(logits, axis=0)
        m_val      = tl.maximum(m_val, block_max)

    # --------------------------------------------------------------------- #
    # Pass-2 : compute lse, softmax, output
    # --------------------------------------------------------------------- #
    lse_denom = tl.zeros((), tl.float32)
    out_acc   = tl.zeros((BLOCK_D, ), tl.float32)

    for t0 in range(0, seq_len, BLOCK_TOKENS):
        tok_offsets = tl.arange(0, BLOCK_TOKENS) + t0
        tok_mask    = tok_offsets < seq_len

        page_ids = tl.load(kv_indices_ptr + tok_offsets, mask=tok_mask, other=0)

        k_ptrs = (
            k_cache_ptr
            + page_ids[:, None] * stride_k_page
            + h_kv             * stride_k_head
            + d_offsets[None, :]
        )
        v_ptrs = (
            v_cache_ptr
            + page_ids[:, None] * stride_v_page
            + h_kv             * stride_v_head
            + d_offsets[None, :]
        )

        k_block = tl.load(k_ptrs, mask=tok_mask[:, None], other=0).to(tl.float32)
        v_block = tl.load(v_ptrs, mask=tok_mask[:, None], other=0).to(tl.float32)

        logits  = tl.sum(k_block * q_vec[None, :], axis=1) * sm_scale
        logits  = logits - m_val
        logits  = tl.where(tok_mask, logits, neg_inf)

        weights = tl.exp(logits)                                        # [T]
        weights = tl.where(tok_mask, weights, 0.0)

        lse_denom += tl.sum(weights, axis=0)

        out_acc += tl.sum(v_block * weights[:, None], axis=0)           # [D]

    out_final = out_acc / lse_denom                                     # [D]

    # --------------------------------------------------------------------- #
    # Write back output and LSE (convert back to bf16 where needed)
    # --------------------------------------------------------------------- #
    out_head_ptr = out_ptr + h_qo * stride_out_head
    tl.store(out_head_ptr + d_offsets, out_final.to(tl.bfloat16))

    lse_head_ptr = lse_ptr + h_qo * stride_lse_head
    INV_LN2      = 1.4426950408889634                                   # 1 / ln(2)
    l_val        = (tl.log(lse_denom) + m_val) * INV_LN2
    tl.store(lse_head_ptr, l_val)


# ----------------------------------------------------------------------------- #
#                                Python Wrapper                                #
# ----------------------------------------------------------------------------- #
@torch.no_grad()
def run(
    q: torch.Tensor,
    k_cache: torch.Tensor,
    v_cache: torch.Tensor,
    kv_indptr: torch.Tensor,
    kv_indices: torch.Tensor,
    sm_scale: float | None = None,
):
    """
    Parameters follow the specification in the problem statement.
    Returns
    -------
    output : bfloat16[batch, 32, 128]
    lse    :  float32[batch, 32]
    """
    if sm_scale is None:
        sm_scale = 1.0 / math.sqrt(128.0)

    # --------------------------------------------------------------------- #
    # Device handling
    # --------------------------------------------------------------------- #
    orig_device = q.device
    cuda_available = torch.cuda.is_available()

    def to_cuda(t):
        if t.device.type == 'cuda':
            return t
        if not cuda_available:
            raise RuntimeError("CUDA is required but not available.")
        return t.cuda(non_blocking=True)

    q          = to_cuda(q).contiguous()
    k_cache    = to_cuda(k_cache).contiguous()
    v_cache    = to_cuda(v_cache).contiguous()
    kv_indptr  = to_cuda(kv_indptr).contiguous()
    kv_indices = to_cuda(kv_indices).contiguous()

    # --------------------------------------------------------------------- #
    # Assertions & shapes
    # --------------------------------------------------------------------- #
    batch_size, num_qo_heads, head_dim = q.shape
    num_pages, page_sz, num_kv_heads, _ = k_cache.shape

    assert num_qo_heads == 32,  "num_qo_heads must be 32"
    assert num_kv_heads == 4,   "num_kv_heads must be 4"
    assert head_dim      == 128
    assert page_sz       == 1
    assert kv_indptr.numel() == batch_size + 1
    assert kv_indices.numel() == kv_indptr[-1].item()

    # Flatten caches (page_size == 1)
    k_cache_flat = k_cache.squeeze(1).contiguous()   # [num_pages, 4, 128]
    v_cache_flat = v_cache.squeeze(1).contiguous()

    # Allocate output tensors
    out = torch.zeros_like(q, dtype=torch.bfloat16, device=q.device)
    lse = torch.full((batch_size, num_qo_heads),
                     -float("inf"), dtype=torch.float32, device=q.device)

    # Strides (in elements, not bytes)
    stride_q_head   = head_dim
    stride_k_head   = head_dim
    stride_k_page   = num_kv_heads * head_dim      # 4 * 128
    stride_v_head   = head_dim
    stride_v_page   = num_kv_heads * head_dim
    stride_out_head = head_dim
    stride_lse_head = 1

    gqa_ratio     = num_qo_heads // num_kv_heads   # 8
    BLOCK_TOKENS  = 128
    BLOCK_D       = 128

    # --------------------------------------------------------------------- #
    # Launch per-sequence kernel
    # --------------------------------------------------------------------- #
    for b in range(batch_size):
        start   = kv_indptr[b].item()
        end     = kv_indptr[b + 1].item()
        seq_len = end - start

        if seq_len == 0:
            continue

        seq_indices = kv_indices[start:end].contiguous()

        grid = (num_qo_heads,)

        gqa_paged_decode_kernel[grid](
            q[b],
            k_cache_flat,
            v_cache_flat,
            seq_indices,
            out[b],
            lse[b],
            sm_scale,
            seq_len,

            stride_q_head,
            stride_k_page,
            stride_k_head,
            stride_v_page,
            stride_v_head,
            stride_out_head,
            stride_lse_head,

            gqa_ratio,
            BLOCK_TOKENS,
            BLOCK_D,

            num_warps = 8,
            num_stages = 4,
        )

    # --------------------------------------------------------------------- #
    # Move outputs back if the original tensors were on CPU
    # --------------------------------------------------------------------- #
    if orig_device.type == 'cpu':
        out = out.to(orig_device)
        lse = lse.to(orig_device)

    return out, lse
scrolls · 253 lines total

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

Best evidence level for this revision: reproducible

JSON