Skip to content
KernelIndex
Search⌘K

gpt-5 / tritoncb1275

gpt-5_triton_cb1275 · gpt-5-2025-08-07 · triton · Apache-2.0

Use it

Vendorable · source mirrored · Apache-2.0View source →

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

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

Benchmark evidence

48 measurements across 1 GPU, fastest first.

Operation / workload
Hardware
Latency
Rank
Observed
GQA paged decode h32 kv8 d128 ps1bf16 · [1, 32, 128] · num_pages=15 · num_kv_indices=14
NVIDIA B200
86.8µs
#7 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [1, 32, 128] · num_pages=67 · num_kv_indices=50
NVIDIA B200
87.4µs
#6 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [1, 32, 128] · num_pages=55 · num_kv_indices=38
NVIDIA B200
87.6µs
#7 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [1, 32, 128] · num_pages=9341 · num_kv_indices=98
NVIDIA B200
88.6µs
#6 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [1, 32, 128] · num_pages=74 · num_kv_indices=57
NVIDIA B200
88.7µs
#6 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [1, 32, 128] · num_pages=82 · num_kv_indices=65
NVIDIA B200
89.3µs
#6 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [1, 32, 128] · num_pages=78 · num_kv_indices=61
NVIDIA B200
89.7µs
#6 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [1, 32, 128] · num_pages=63 · num_kv_indices=46
NVIDIA B200
92.6µs
#7 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [1, 32, 128] · num_pages=59 · num_kv_indices=42
NVIDIA B200
92.7µs
#7 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [1, 32, 128] · num_pages=11 · num_kv_indices=10
NVIDIA B200
93.8µs
#7 of 7
2025-10-16
Show all 48 measurements ›
GQA paged decode h32 kv8 d128 ps1bf16 · [1, 32, 128] · num_pages=71 · num_kv_indices=54
NVIDIA B200
94.4µs
#6 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [1, 32, 128] · num_pages=18 · num_kv_indices=2
NVIDIA B200
94.4µs
#7 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [1, 32, 128] · num_pages=9316 · num_kv_indices=73
NVIDIA B200
94.9µs
#6 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [1, 32, 128] · num_pages=406 · num_kv_indices=356
NVIDIA B200
102.0µs
#5 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [16, 32, 128] · num_pages=1069 · num_kv_indices=1034
NVIDIA B200
112.7µs
#5 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [16, 32, 128] · num_pages=1220 · num_kv_indices=1193
NVIDIA B200
113.2µs
#5 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [1, 32, 128] · num_pages=597 · num_kv_indices=547
NVIDIA B200
113.3µs
#5 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [16, 32, 128] · num_pages=1396 · num_kv_indices=1369
NVIDIA B200
114.0µs
#5 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [16, 32, 128] · num_pages=1556 · num_kv_indices=1529
NVIDIA B200
114.3µs
#5 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [16, 32, 128] · num_pages=2052 · num_kv_indices=2025
NVIDIA B200
114.3µs
#5 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [1, 32, 128] · num_pages=223 · num_kv_indices=173
NVIDIA B200
114.4µs
#5 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [16, 32, 128] · num_pages=1732 · num_kv_indices=1705
NVIDIA B200
115.5µs
#5 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [16, 32, 128] · num_pages=1892 · num_kv_indices=1865
NVIDIA B200
116.3µs
#5 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [16, 32, 128] · num_pages=2212 · num_kv_indices=2185
NVIDIA B200
117.8µs
#5 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [16, 32, 128] · num_pages=2548 · num_kv_indices=2521
NVIDIA B200
122.7µs
#5 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [16, 32, 128] · num_pages=2372 · num_kv_indices=2345
NVIDIA B200
123.6µs
#5 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [16, 32, 128] · num_pages=3044 · num_kv_indices=3017
NVIDIA B200
123.7µs
#5 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [16, 32, 128] · num_pages=2868 · num_kv_indices=2841
NVIDIA B200
124.9µs
#5 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [16, 32, 128] · num_pages=2708 · num_kv_indices=2681
NVIDIA B200
125.0µs
#5 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [16, 32, 128] · num_pages=4333 · num_kv_indices=4298
NVIDIA B200
133.3µs
#5 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [16, 32, 128] · num_pages=81390 · num_kv_indices=12942
NVIDIA B200
212.5µs
#5 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [16, 32, 128] · num_pages=30163 · num_kv_indices=20911
NVIDIA B200
250.2µs
#5 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [64, 32, 128] · num_pages=60071 · num_kv_indices=50902
NVIDIA B200
435.7µs
#5 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [64, 32, 128] · num_pages=62605 · num_kv_indices=53334
NVIDIA B200
437.9µs
#5 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [64, 32, 128] · num_pages=63053 · num_kv_indices=53782
NVIDIA B200
441.4µs
#5 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [64, 32, 128] · num_pages=63821 · num_kv_indices=54550
NVIDIA B200
444.7µs
#4 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [64, 32, 128] · num_pages=65037 · num_kv_indices=55766
NVIDIA B200
444.7µs
#5 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [64, 32, 128] · num_pages=66253 · num_kv_indices=56982
NVIDIA B200
447.5µs
#5 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [64, 32, 128] · num_pages=64205 · num_kv_indices=54934
NVIDIA B200
447.6µs
#5 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [64, 32, 128] · num_pages=64653 · num_kv_indices=55382
NVIDIA B200
448.0µs
#4 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [64, 32, 128] · num_pages=67021 · num_kv_indices=57750
NVIDIA B200
448.2µs
#5 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [64, 32, 128] · num_pages=65805 · num_kv_indices=56534
NVIDIA B200
450.6µs
#5 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [64, 32, 128] · num_pages=66637 · num_kv_indices=57366
NVIDIA B200
452.2µs
#4 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [64, 32, 128] · num_pages=65421 · num_kv_indices=56150
NVIDIA B200
452.5µs
#5 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [64, 32, 128] · num_pages=67405 · num_kv_indices=58134
NVIDIA B200
458.5µs
#5 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [64, 32, 128] · num_pages=67853 · num_kv_indices=58582
NVIDIA B200
462.7µs
#5 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [64, 32, 128] · num_pages=68237 · num_kv_indices=58966
NVIDIA B200
467.9µs
#5 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [64, 32, 128] · num_pages=63437 · num_kv_indices=54166
NVIDIA B200
469.2µs
#5 of 7
2025-10-16

Reported · How evidence levels are derived →

Source and license

sourcehttps://huggingface.co/datasets/flashinfer-ai/flashinfer-trace
commitda915083d4c7c5e61aa3005e3d17ae488e0fc71c
revision digestsha256:10231abe2f3401b0219cd1d797d674a685a6d3a88856dbfa7eb880b115a36a37
license declaredApache-2.0
license concludedApache-2.0
authorsgpt-5-2025-08-07
imported2026-08-20

Techniques

Extracted from the mirrored source by pattern, never inferred. Each row cites its line.

num-warps = 8num_warps=8, num_stages=3,
stages = 3num_warps=8, num_stages=3,

Kernel source

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


@triton.jit
def gqa_paged_decode_h32_kv8_d128_ps1_kernel(
    q_ptr, k_ptr, v_ptr,
    kv_indptr_ptr, kv_indices_ptr,
    out_ptr, lse_ptr,
    sm_scale_ptr, inv_ln2_ptr,
    batch_size,
    stride_q_b, stride_q_h, stride_q_d,
    stride_k_p, stride_k_ps, stride_k_h, stride_k_d,
    stride_v_p, stride_v_ps, stride_v_h, stride_v_d,
    stride_out_b, stride_out_h, stride_out_d,
    stride_lse_b, stride_lse_h,
    BLOCK_T: tl.constexpr, BLOCK_D: tl.constexpr, STEP: tl.constexpr,
    GQA_RATIO: tl.constexpr,
):
    pid = tl.program_id(0)
    num_qo_heads = 32
    b = pid // num_qo_heads
    h = pid % num_qo_heads
    if b >= batch_size:
        return

    # Load scalar parameters
    sm_scale = tl.load(sm_scale_ptr)
    inv_ln2 = tl.load(inv_ln2_ptr)

    b_i64 = b.to(tl.int64)
    h_i64 = h.to(tl.int64)

    # GQA mapping
    kv_head = (h // GQA_RATIO)
    kv_head_i64 = kv_head.to(tl.int64)

    # Load start/end pointers for this batch element
    page_start = tl.load(kv_indptr_ptr + b_i64)
    page_end = tl.load(kv_indptr_ptr + b_i64 + 1)
    n_tokens = page_end - page_start

    # Prepare output/lse pointers
    d_all = tl.arange(0, BLOCK_D)
    out_row_ptrs = out_ptr + b_i64 * stride_out_b + h_i64 * stride_out_h + d_all.to(tl.int64) * stride_out_d
    lse_ptr_ = lse_ptr + b_i64 * stride_lse_b + h_i64 * stride_lse_h

    # If no tokens, write zeros and -inf LSE and return
    if n_tokens <= 0:
        zero_bf16 = tl.zeros([BLOCK_D], dtype=tl.bfloat16)
        tl.store(out_row_ptrs, zero_bf16, mask=d_all < BLOCK_D)
        tl.store(lse_ptr_, -float("inf"))
        return

    # Preload Q in four STEP chunks (bf16 -> fp32)
    # chunk 0
    d_idx0 = tl.arange(0, STEP)
    q0 = tl.load(
        q_ptr + b_i64 * stride_q_b + h_i64 * stride_q_h + d_idx0.to(tl.int64) * stride_q_d,
        mask=d_idx0 < BLOCK_D,
        other=0,
    ).to(tl.float32)
    # chunk 1
    d_idx1 = STEP + tl.arange(0, STEP)
    q1 = tl.load(
        q_ptr + b_i64 * stride_q_b + h_i64 * stride_q_h + d_idx1.to(tl.int64) * stride_q_d,
        mask=d_idx1 < BLOCK_D,
        other=0,
    ).to(tl.float32)
    # chunk 2
    d_idx2 = (2 * STEP) + tl.arange(0, STEP)
    q2 = tl.load(
        q_ptr + b_i64 * stride_q_b + h_i64 * stride_q_h + d_idx2.to(tl.int64) * stride_q_d,
        mask=d_idx2 < BLOCK_D,
        other=0,
    ).to(tl.float32)
    # chunk 3
    d_idx3 = (3 * STEP) + tl.arange(0, STEP)
    q3 = tl.load(
        q_ptr + b_i64 * stride_q_b + h_i64 * stride_q_h + d_idx3.to(tl.int64) * stride_q_d,
        mask=d_idx3 < BLOCK_D,
        other=0,
    ).to(tl.float32)

    # Streaming softmax variables
    m = -float("inf")
    l = 0.0

    # Accumulator for output across head_dim in 4 chunks (STEP each)
    out_acc0 = tl.zeros([STEP], dtype=tl.float32)
    out_acc1 = tl.zeros([STEP], dtype=tl.float32)
    out_acc2 = tl.zeros([STEP], dtype=tl.float32)
    out_acc3 = tl.zeros([STEP], dtype=tl.float32)

    # Loop over tokens in blocks of BLOCK_T
    pos = 0
    while pos < n_tokens:
        t_offsets = tl.arange(0, BLOCK_T)
        offs = pos + t_offsets
        mask_t = offs < n_tokens

        # Gather page indices for this block
        idx = tl.load(kv_indices_ptr + page_start.to(tl.int64) + offs.to(tl.int64), mask=mask_t, other=0)

        # Compute logits for this block: [BLOCK_T]
        logits = tl.zeros([BLOCK_T], dtype=tl.float32)

        # chunk 0
        k_ptrs0 = (
            k_ptr
            + (idx[:, None].to(tl.int64) * stride_k_p)
            + (kv_head_i64 * stride_k_h)
            + (d_idx0[None, :].to(tl.int64) * stride_k_d)
        )
        k0 = tl.load(k_ptrs0, mask=mask_t[:, None] & (d_idx0[None, :] < BLOCK_D), other=0).to(tl.float32)
        logits += tl.sum(k0 * q0[None, :], axis=1)

        # chunk 1
        k_ptrs1 = (
            k_ptr
            + (idx[:, None].to(tl.int64) * stride_k_p)
            + (kv_head_i64 * stride_k_h)
            + (d_idx1[None, :].to(tl.int64) * stride_k_d)
        )
        k1 = tl.load(k_ptrs1, mask=mask_t[:, None] & (d_idx1[None, :] < BLOCK_D), other=0).to(tl.float32)
        logits += tl.sum(k1 * q1[None, :], axis=1)

        # chunk 2
        k_ptrs2 = (
            k_ptr
            + (idx[:, None].to(tl.int64) * stride_k_p)
            + (kv_head_i64 * stride_k_h)
            + (d_idx2[None, :].to(tl.int64) * stride_k_d)
        )
        k2 = tl.load(k_ptrs2, mask=mask_t[:, None] & (d_idx2[None, :] < BLOCK_D), other=0).to(tl.float32)
        logits += tl.sum(k2 * q2[None, :], axis=1)

        # chunk 3
        k_ptrs3 = (
            k_ptr
            + (idx[:, None].to(tl.int64) * stride_k_p)
            + (kv_head_i64 * stride_k_h)
            + (d_idx3[None, :].to(tl.int64) * stride_k_d)
        )
        k3 = tl.load(k_ptrs3, mask=mask_t[:, None] & (d_idx3[None, :] < BLOCK_D), other=0).to(tl.float32)
        logits += tl.sum(k3 * q3[None, :], axis=1)

        # Scale logits and apply mask
        logits = logits * sm_scale
        logits = tl.where(mask_t, logits, -float("inf"))

        # Compute block max and update running m and l
        block_max = tl.max(logits, axis=0)
        new_m = tl.maximum(m, block_max)
        scale_old = tl.exp(m - new_m)

        # Weights for this block
        weights = tl.exp(logits - new_m)

        # Update l
        l = l * scale_old + tl.sum(weights, axis=0)

        # Scale previous accumulators by scale_old
        out_acc0 *= scale_old
        out_acc1 *= scale_old
        out_acc2 *= scale_old
        out_acc3 *= scale_old

        # Accumulate V weighted by weights
        # chunk 0
        v_ptrs0 = (
            v_ptr
            + (idx[:, None].to(tl.int64) * stride_v_p)
            + (kv_head_i64 * stride_v_h)
            + (d_idx0[None, :].to(tl.int64) * stride_v_d)
        )
        v0 = tl.load(v_ptrs0, mask=mask_t[:, None] & (d_idx0[None, :] < BLOCK_D), other=0).to(tl.float32)
        out_acc0 += tl.sum(v0 * weights[:, None], axis=0)

        # chunk 1
        v_ptrs1 = (
            v_ptr
            + (idx[:, None].to(tl.int64) * stride_v_p)
            + (kv_head_i64 * stride_v_h)
            + (d_idx1[None, :].to(tl.int64) * stride_v_d)
        )
        v1 = tl.load(v_ptrs1, mask=mask_t[:, None] & (d_idx1[None, :] < BLOCK_D), other=0).to(tl.float32)
        out_acc1 += tl.sum(v1 * weights[:, None], axis=0)

        # chunk 2
        v_ptrs2 = (
            v_ptr
            + (idx[:, None].to(tl.int64) * stride_v_p)
            + (kv_head_i64 * stride_v_h)
            + (d_idx2[None, :].to(tl.int64) * stride_v_d)
        )
        v2 = tl.load(v_ptrs2, mask=mask_t[:, None] & (d_idx2[None, :] < BLOCK_D), other=0).to(tl.float32)
        out_acc2 += tl.sum(v2 * weights[:, None], axis=0)

        # chunk 3
        v_ptrs3 = (
            v_ptr
            + (idx[:, None].to(tl.int64) * stride_v_p)
            + (kv_head_i64 * stride_v_h)
            + (d_idx3[None, :].to(tl.int64) * stride_v_d)
        )
        v3 = tl.load(v_ptrs3, mask=mask_t[:, None] & (d_idx3[None, :] < BLOCK_D), other=0).to(tl.float32)
        out_acc3 += tl.sum(v3 * weights[:, None], axis=0)

        # Update running max
        m = new_m
        pos += BLOCK_T

    # Finalize lse in base 2
    lse_base2 = (tl.log(l) + m) * inv_ln2
    tl.store(lse_ptr_, lse_base2)

    # Normalize by l
    inv_l = 1.0 / l
    out_acc0 *= inv_l
    out_acc1 *= inv_l
    out_acc2 *= inv_l
    out_acc3 *= inv_l

    # Store output chunks
    # chunk 0
    tl.store(
        out_ptr + b_i64 * stride_out_b + h_i64 * stride_out_h + (tl.arange(0, STEP).to(tl.int64)) * stride_out_d,
        out_acc0.to(tl.bfloat16),
        mask=(tl.arange(0, STEP) < BLOCK_D),
    )
    # chunk 1
    tl.store(
        out_ptr + b_i64 * stride_out_b + h_i64 * stride_out_h + ((STEP + tl.arange(0, STEP)).to(tl.int64)) * stride_out_d,
        out_acc1.to(tl.bfloat16),
        mask=((STEP + tl.arange(0, STEP)) < BLOCK_D),
    )
    # chunk 2
    tl.store(
        out_ptr + b_i64 * stride_out_b + h_i64 * stride_out_h + (((2 * STEP) + tl.arange(0, STEP)).to(tl.int64)) * stride_out_d,
        out_acc2.to(tl.bfloat16),
        mask=(((2 * STEP) + tl.arange(0, STEP)) < BLOCK_D),
    )
    # chunk 3
    tl.store(
        out_ptr + b_i64 * stride_out_b + h_i64 * stride_out_h + (((3 * STEP) + tl.arange(0, STEP)).to(tl.int64)) * stride_out_d,
        out_acc3.to(tl.bfloat16),
        mask=(((3 * STEP) + tl.arange(0, STEP)) < BLOCK_D),
    )


def run(q, k_cache, v_cache, kv_indptr, kv_indices, sm_scale=None):
    # Validate inputs and move to CUDA if available
    if not torch.cuda.is_available():
        # If any tensor is already on CUDA but CUDA is unavailable, raise error
        if any(t.is_cuda for t in [q, k_cache, v_cache, kv_indptr, kv_indices] if isinstance(t, torch.Tensor)):
            raise RuntimeError("CUDA is not available but some inputs are CUDA tensors.")
        raise RuntimeError("CUDA is required to run Triton kernels. Please enable a CUDA-capable device.")

    device_out = q.device

    def to_cuda(t):
        return t if t.is_cuda else t.cuda()

    q_c = to_cuda(q)
    k_c = to_cuda(k_cache)
    v_c = to_cuda(v_cache)
    kv_indptr_c = to_cuda(kv_indptr)
    kv_indices_c = to_cuda(kv_indices)

    # Check dtypes and shapes
    assert q_c.dtype == torch.bfloat16, "q must be bfloat16"
    assert k_c.dtype == torch.bfloat16 and v_c.dtype == torch.bfloat16, "k_cache and v_cache must be bfloat16"
    assert kv_indptr_c.dtype == torch.int32, "kv_indptr must be int32"
    assert kv_indices_c.dtype == torch.int32, "kv_indices must be int32"

    batch_size, num_qo_heads, head_dim = q_c.shape
    num_pages, page_size, num_kv_heads, head_dim_k = k_c.shape
    assert num_qo_heads == 32, "num_qo_heads must be 32"
    assert num_kv_heads == 8, "num_kv_heads must be 8"
    assert head_dim == 128 and head_dim_k == 128, "head_dim must be 128"
    assert page_size == 1, "page_size must be 1"

    len_indptr = kv_indptr_c.shape[0]
    num_kv_indices = kv_indices_c.shape[0]
    assert len_indptr == batch_size + 1, "len_indptr must equal batch_size + 1"
    last = kv_indptr_c[-1].item()
    assert num_kv_indices == last, "num_kv_indices must equal kv_indptr[-1].item()"

    # Default softmax scale
    if sm_scale is None:
        sm_scale = 1.0 / math.sqrt(head_dim)
    if not isinstance(sm_scale, torch.Tensor):
        sm_scale_t = torch.tensor(sm_scale, dtype=torch.float32, device=q_c.device)
    else:
        sm_scale_t = sm_scale.to(dtype=torch.float32, device=q_c.device)

    inv_ln2 = torch.tensor(1.0 / math.log(2.0), dtype=torch.float32, device=q_c.device)

    # Allocate outputs
    output_c = torch.empty((batch_size, num_qo_heads, head_dim), dtype=torch.bfloat16, device=q_c.device)
    lse_c = torch.empty((batch_size, num_qo_heads), dtype=torch.float32, device=q_c.device)

    # Extract strides
    stride_q_b, stride_q_h, stride_q_d = q_c.stride()
    stride_k_p, stride_k_ps, stride_k_h, stride_k_d = k_c.stride()
    stride_v_p, stride_v_ps, stride_v_h, stride_v_d = v_c.stride()
    stride_out_b, stride_out_h, stride_out_d = output_c.stride()
    stride_lse_b, stride_lse_h = lse_c.stride()

    # Launch kernel
    BLOCK_D = 128
    STEP = 32
    BLOCK_T = 128
    GQA_RATIO = 4

    grid = (batch_size * num_qo_heads,)

    gqa_paged_decode_h32_kv8_d128_ps1_kernel[grid](
        q_c, k_c, v_c,
        kv_indptr_c, kv_indices_c,
        output_c, lse_c,
        sm_scale_t, inv_ln2,
        batch_size,
        stride_q_b, stride_q_h, stride_q_d,
        stride_k_p, stride_k_ps, stride_k_h, stride_k_d,
        stride_v_p, stride_v_ps, stride_v_h, stride_v_d,
        stride_out_b, stride_out_h, stride_out_d,
        stride_lse_b, stride_lse_h,
        BLOCK_T=BLOCK_T, BLOCK_D=BLOCK_D, STEP=STEP, GQA_RATIO=GQA_RATIO,
        num_warps=8, num_stages=3,
    )

    # Move outputs back to original device of q if needed
    if output_c.device != device_out:
        output = output_c.to(device_out)
        lse = lse_c.to(device_out)
    else:
        output = output_c
        lse = lse_c

    return output, lse
scrolls · 344 lines total

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

Best evidence level for this revision: reported

JSON