Skip to content
KernelIndex
Search⌘K

gpt-5 / triton7308c5

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

Use it

Vendorable · source mirrored · Apache-2.0View source →

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

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

Benchmark evidence

21 measurements across 1 GPU, fastest first.

Operation / workload
Hardware
Latency
Rank
Observed
GQA ragged prefill causal h32 kv8 d128bf16 · [7, 8, 128] · #cb4b00
NVIDIA B200
146.3µs
#5 of 5
2025-10-19
NVIDIA B200
146.7µs
#9 of 10
2025-10-19
GQA ragged prefill causal h32 kv8 d128bf16 · [7, 8, 128] · #227ae2
NVIDIA B200
147.3µs
#8 of 10
2025-10-19
GQA ragged prefill causal h32 kv8 d128bf16 · [1, 8, 128] · #5fa6fa
NVIDIA B200
147.4µs
#5 of 5
2025-10-19
GQA ragged prefill causal h32 kv8 d128bf16 · [1, 8, 128] · #162cf6
NVIDIA B200
148.1µs
#9 of 10
2025-10-19
GQA ragged prefill causal h32 kv8 d128bf16 · [7, 8, 128] · #227ae2
NVIDIA B200
148.3µs
#9 of 10
2025-10-19
GQA ragged prefill causal h32 kv8 d128bf16 · [35, 8, 128] · #f3d59b
NVIDIA B200
148.8µs
#3 of 5
2025-10-19
GQA ragged prefill causal h32 kv8 d128bf16 · [35, 8, 128] · #10e83c
NVIDIA B200
149.0µs
#9 of 20
2025-10-19
GQA ragged prefill causal h32 kv8 d128bf16 · [1, 8, 128] · #162cf6
NVIDIA B200
149.9µs
#10 of 10
2025-10-19
GQA ragged prefill causal h32 kv8 d128bf16 · [35, 8, 128] · #10e83c
NVIDIA B200
150.0µs
#10 of 20
2025-10-19
Show all 21 measurements ›
GQA ragged prefill causal h32 kv8 d128bf16 · [35, 8, 128] · #10e83c
NVIDIA B200
150.2µs
#11 of 20
2025-10-19
GQA ragged prefill causal h32 kv8 d128bf16 · [35, 8, 128] · #10e83c
NVIDIA B200
150.5µs
#12 of 20
2025-10-19
NVIDIA B200
151.8µs
#10 of 10
2025-10-19
NVIDIA B200
175.7µs
#5 of 10
2025-10-19
NVIDIA B200
176.1µs
#6 of 10
2025-10-19
NVIDIA B200
193.5µs
#3 of 5
2025-10-19
NVIDIA B200
245.1µs
#4 of 5
2025-10-19
NVIDIA B200
1.10ms
#2 of 5
2025-10-19
NVIDIA B200
51.8ms
#2 of 5
2025-10-19
NVIDIA B200
51.9ms
#2 of 5
2025-10-19
NVIDIA B200
52.8ms
#2 of 5
2025-10-19

Reported · How evidence levels are derived →

Source and license

sourcehttps://huggingface.co/datasets/flashinfer-ai/flashinfer-trace
commitda915083d4c7c5e61aa3005e3d17ae488e0fc71c
revision digestsha256:561ea20168ac707b05662f1c596a5f16c3ebf9160e26614c28d99fe0601509a5
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 = 4num_warps=4, num_stages=2,
stages = 2num_warps=4, num_stages=2,
tile-n = 64BLOCK_N = 64

Kernel source

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


@triton.jit
def gqa_ragged_prefill_causal_h32_kv8_d128_kernel(
    q_ptr, k_ptr, v_ptr,
    stride_q_q, stride_q_h, stride_q_d,
    stride_k_k, stride_k_h, stride_k_d,
    stride_v_k, stride_v_h, stride_v_d,
    out_ptr, stride_out_q, stride_out_h, stride_out_d,
    lse_ptr, stride_lse_q, stride_lse_h,
    q_kv_start_ptr, q_kv_max_ptr,
    total_q,
    sm_scale, ln2,
    RATIO: tl.constexpr, HEAD_DIM: tl.constexpr,
    BLOCK_N: tl.constexpr, BLOCK_DK: tl.constexpr, BLOCK_DV: tl.constexpr,
):
    pid_q = tl.program_id(0)
    kvh = tl.program_id(1)
    if pid_q >= total_q:
        return

    kv_start = tl.load(q_kv_start_ptr + pid_q, mask=True, other=0).to(tl.int32)
    kv_max = tl.load(q_kv_max_ptr + pid_q, mask=True, other=0).to(tl.int32)

    heads_base = kvh * RATIO
    neg_inf = tl.full([], -float("inf"), tl.float32)

    if kv_max <= 0:
        # No available keys for this query; set LSE to -inf and outputs to 0
        for r in range(RATIO):
            lse_ptr_r = lse_ptr + pid_q * stride_lse_q + (heads_base + r) * stride_lse_h
            tl.store(lse_ptr_r, neg_inf)
            # store output zeros
            for dv0 in range(0, HEAD_DIM, BLOCK_DV):
                d_voffs = dv0 + tl.arange(0, BLOCK_DV)
                out_ptrs = out_ptr + pid_q * stride_out_q + (heads_base + r) * stride_out_h + d_voffs * stride_out_d
                tl.store(out_ptrs, tl.zeros([BLOCK_DV], dtype=tl.bfloat16))
        return

    # Initialize streaming softmax stats per head (RATIO=4)
    m0 = neg_inf
    m1 = neg_inf
    m2 = neg_inf
    m3 = neg_inf
    l0 = tl.zeros([], dtype=tl.float32)
    l1 = tl.zeros([], dtype=tl.float32)
    l2 = tl.zeros([], dtype=tl.float32)
    l3 = tl.zeros([], dtype=tl.float32)

    # Output accumulators per head, split along D into 4 segments (HEAD_DIM=128, BLOCK_DV=32)
    o0_s0 = tl.zeros([BLOCK_DV], dtype=tl.float32)
    o0_s1 = tl.zeros([BLOCK_DV], dtype=tl.float32)
    o0_s2 = tl.zeros([BLOCK_DV], dtype=tl.float32)
    o0_s3 = tl.zeros([BLOCK_DV], dtype=tl.float32)

    o1_s0 = tl.zeros([BLOCK_DV], dtype=tl.float32)
    o1_s1 = tl.zeros([BLOCK_DV], dtype=tl.float32)
    o1_s2 = tl.zeros([BLOCK_DV], dtype=tl.float32)
    o1_s3 = tl.zeros([BLOCK_DV], dtype=tl.float32)

    o2_s0 = tl.zeros([BLOCK_DV], dtype=tl.float32)
    o2_s1 = tl.zeros([BLOCK_DV], dtype=tl.float32)
    o2_s2 = tl.zeros([BLOCK_DV], dtype=tl.float32)
    o2_s3 = tl.zeros([BLOCK_DV], dtype=tl.float32)

    o3_s0 = tl.zeros([BLOCK_DV], dtype=tl.float32)
    o3_s1 = tl.zeros([BLOCK_DV], dtype=tl.float32)
    o3_s2 = tl.zeros([BLOCK_DV], dtype=tl.float32)
    o3_s3 = tl.zeros([BLOCK_DV], dtype=tl.float32)

    # Loop over key tiles
    for start_n in range(0, kv_max, BLOCK_N):
        key_offsets = start_n + tl.arange(0, BLOCK_N)
        key_mask = key_offsets < kv_max

        # Accumulate logits per head for this tile
        logits0 = tl.zeros([BLOCK_N], dtype=tl.float32)
        logits1 = tl.zeros([BLOCK_N], dtype=tl.float32)
        logits2 = tl.zeros([BLOCK_N], dtype=tl.float32)
        logits3 = tl.zeros([BLOCK_N], dtype=tl.float32)

        for d0 in range(0, HEAD_DIM, BLOCK_DK):
            d_off = d0 + tl.arange(0, BLOCK_DK)

            # Load K chunk: [BLOCK_N, BLOCK_DK] -> fp32
            k_ptrs = k_ptr + (kv_start + key_offsets)[:, None] * stride_k_k + kvh * stride_k_h + d_off[None, :] * stride_k_d
            k_chunk = tl.load(
                k_ptrs,
                mask=key_mask[:, None] & (d_off[None, :] < HEAD_DIM),
                other=0
            ).to(tl.float32)

            # Load Q chunk and accumulate logits for each of the 4 query heads in this kv head group
            # Head 0
            q_ptrs0 = q_ptr + pid_q * stride_q_q + (heads_base + 0) * stride_q_h + d_off * stride_q_d
            q_vec0 = tl.load(q_ptrs0, mask=d_off < HEAD_DIM, other=0).to(tl.float32)
            logits0 += tl.sum(k_chunk * q_vec0[None, :], axis=1)

            # Head 1
            q_ptrs1 = q_ptr + pid_q * stride_q_q + (heads_base + 1) * stride_q_h + d_off * stride_q_d
            q_vec1 = tl.load(q_ptrs1, mask=d_off < HEAD_DIM, other=0).to(tl.float32)
            logits1 += tl.sum(k_chunk * q_vec1[None, :], axis=1)

            # Head 2
            q_ptrs2 = q_ptr + pid_q * stride_q_q + (heads_base + 2) * stride_q_h + d_off * stride_q_d
            q_vec2 = tl.load(q_ptrs2, mask=d_off < HEAD_DIM, other=0).to(tl.float32)
            logits2 += tl.sum(k_chunk * q_vec2[None, :], axis=1)

            # Head 3
            q_ptrs3 = q_ptr + pid_q * stride_q_q + (heads_base + 3) * stride_q_h + d_off * stride_q_d
            q_vec3 = tl.load(q_ptrs3, mask=d_off < HEAD_DIM, other=0).to(tl.float32)
            logits3 += tl.sum(k_chunk * q_vec3[None, :], axis=1)

        # Scale and apply mask
        p0 = logits0 * sm_scale
        p1 = logits1 * sm_scale
        p2 = logits2 * sm_scale
        p3 = logits3 * sm_scale

        p0 = tl.where(key_mask, p0, neg_inf)
        p1 = tl.where(key_mask, p1, neg_inf)
        p2 = tl.where(key_mask, p2, neg_inf)
        p3 = tl.where(key_mask, p3, neg_inf)

        # Preload V chunks once for the tile (reused across heads)
        d0_idx = 0 + tl.arange(0, BLOCK_DV)
        d1_idx = BLOCK_DV + tl.arange(0, BLOCK_DV)
        d2_idx = 2 * BLOCK_DV + tl.arange(0, BLOCK_DV)
        d3_idx = 3 * BLOCK_DV + tl.arange(0, BLOCK_DV)

        v_ptrs0 = v_ptr + (kv_start + key_offsets)[:, None] * stride_v_k + kvh * stride_v_h + d0_idx[None, :] * stride_v_d
        v_ptrs1 = v_ptr + (kv_start + key_offsets)[:, None] * stride_v_k + kvh * stride_v_h + d1_idx[None, :] * stride_v_d
        v_ptrs2 = v_ptr + (kv_start + key_offsets)[:, None] * stride_v_k + kvh * stride_v_h + d2_idx[None, :] * stride_v_d
        v_ptrs3 = v_ptr + (kv_start + key_offsets)[:, None] * stride_v_k + kvh * stride_v_h + d3_idx[None, :] * stride_v_d

        v_chunk0 = tl.load(v_ptrs0, mask=key_mask[:, None], other=0).to(tl.float32)
        v_chunk1 = tl.load(v_ptrs1, mask=key_mask[:, None], other=0).to(tl.float32)
        v_chunk2 = tl.load(v_ptrs2, mask=key_mask[:, None], other=0).to(tl.float32)
        v_chunk3 = tl.load(v_ptrs3, mask=key_mask[:, None], other=0).to(tl.float32)

        # Head 0
        m0_tile = tl.max(p0, axis=0)
        m0_new = tl.maximum(m0, m0_tile)
        alpha0 = tl.exp(m0 - m0_new)
        o0_s0 = o0_s0 * alpha0
        o0_s1 = o0_s1 * alpha0
        o0_s2 = o0_s2 * alpha0
        o0_s3 = o0_s3 * alpha0
        w0 = tl.exp(p0 - m0_new)
        l0 = l0 * alpha0 + tl.sum(w0, axis=0)
        o0_s0 = o0_s0 + tl.sum(v_chunk0 * w0[:, None], axis=0)
        o0_s1 = o0_s1 + tl.sum(v_chunk1 * w0[:, None], axis=0)
        o0_s2 = o0_s2 + tl.sum(v_chunk2 * w0[:, None], axis=0)
        o0_s3 = o0_s3 + tl.sum(v_chunk3 * w0[:, None], axis=0)
        m0 = m0_new

        # Head 1
        m1_tile = tl.max(p1, axis=0)
        m1_new = tl.maximum(m1, m1_tile)
        alpha1 = tl.exp(m1 - m1_new)
        o1_s0 = o1_s0 * alpha1
        o1_s1 = o1_s1 * alpha1
        o1_s2 = o1_s2 * alpha1
        o1_s3 = o1_s3 * alpha1
        w1 = tl.exp(p1 - m1_new)
        l1 = l1 * alpha1 + tl.sum(w1, axis=0)
        o1_s0 = o1_s0 + tl.sum(v_chunk0 * w1[:, None], axis=0)
        o1_s1 = o1_s1 + tl.sum(v_chunk1 * w1[:, None], axis=0)
        o1_s2 = o1_s2 + tl.sum(v_chunk2 * w1[:, None], axis=0)
        o1_s3 = o1_s3 + tl.sum(v_chunk3 * w1[:, None], axis=0)
        m1 = m1_new

        # Head 2
        m2_tile = tl.max(p2, axis=0)
        m2_new = tl.maximum(m2, m2_tile)
        alpha2 = tl.exp(m2 - m2_new)
        o2_s0 = o2_s0 * alpha2
        o2_s1 = o2_s1 * alpha2
        o2_s2 = o2_s2 * alpha2
        o2_s3 = o2_s3 * alpha2
        w2 = tl.exp(p2 - m2_new)
        l2 = l2 * alpha2 + tl.sum(w2, axis=0)
        o2_s0 = o2_s0 + tl.sum(v_chunk0 * w2[:, None], axis=0)
        o2_s1 = o2_s1 + tl.sum(v_chunk1 * w2[:, None], axis=0)
        o2_s2 = o2_s2 + tl.sum(v_chunk2 * w2[:, None], axis=0)
        o2_s3 = o2_s3 + tl.sum(v_chunk3 * w2[:, None], axis=0)
        m2 = m2_new

        # Head 3
        m3_tile = tl.max(p3, axis=0)
        m3_new = tl.maximum(m3, m3_tile)
        alpha3 = tl.exp(m3 - m3_new)
        o3_s0 = o3_s0 * alpha3
        o3_s1 = o3_s1 * alpha3
        o3_s2 = o3_s2 * alpha3
        o3_s3 = o3_s3 * alpha3
        w3 = tl.exp(p3 - m3_new)
        l3 = l3 * alpha3 + tl.sum(w3, axis=0)
        o3_s0 = o3_s0 + tl.sum(v_chunk0 * w3[:, None], axis=0)
        o3_s1 = o3_s1 + tl.sum(v_chunk1 * w3[:, None], axis=0)
        o3_s2 = o3_s2 + tl.sum(v_chunk2 * w3[:, None], axis=0)
        o3_s3 = o3_s3 + tl.sum(v_chunk3 * w3[:, None], axis=0)
        m3 = m3_new

    # Finalize: compute output = O / l, and lse = (m + log(l)) / ln2
    d0 = 0 + tl.arange(0, BLOCK_DV)
    d1 = BLOCK_DV + tl.arange(0, BLOCK_DV)
    d2 = 2 * BLOCK_DV + tl.arange(0, BLOCK_DV)
    d3 = 3 * BLOCK_DV + tl.arange(0, BLOCK_DV)

    # Head 0
    l0_pos = l0 > 0
    lse0 = tl.where(l0_pos, (m0 + tl.log(l0)) / ln2, neg_inf)
    lse_ptr0 = lse_ptr + pid_q * stride_lse_q + (heads_base + 0) * stride_lse_h
    tl.store(lse_ptr0, lse0)
    out_ptrs0 = out_ptr + pid_q * stride_out_q + (heads_base + 0) * stride_out_h
    o0_s0_out = tl.where(l0_pos, o0_s0 / l0, tl.zeros([BLOCK_DV], dtype=tl.float32))
    o0_s1_out = tl.where(l0_pos, o0_s1 / l0, tl.zeros([BLOCK_DV], dtype=tl.float32))
    o0_s2_out = tl.where(l0_pos, o0_s2 / l0, tl.zeros([BLOCK_DV], dtype=tl.float32))
    o0_s3_out = tl.where(l0_pos, o0_s3 / l0, tl.zeros([BLOCK_DV], dtype=tl.float32))
    tl.store(out_ptrs0 + d0 * stride_out_d, o0_s0_out.to(tl.bfloat16))
    tl.store(out_ptrs0 + d1 * stride_out_d, o0_s1_out.to(tl.bfloat16))
    tl.store(out_ptrs0 + d2 * stride_out_d, o0_s2_out.to(tl.bfloat16))
    tl.store(out_ptrs0 + d3 * stride_out_d, o0_s3_out.to(tl.bfloat16))

    # Head 1
    l1_pos = l1 > 0
    lse1 = tl.where(l1_pos, (m1 + tl.log(l1)) / ln2, neg_inf)
    lse_ptr1 = lse_ptr + pid_q * stride_lse_q + (heads_base + 1) * stride_lse_h
    tl.store(lse_ptr1, lse1)
    out_ptrs1 = out_ptr + pid_q * stride_out_q + (heads_base + 1) * stride_out_h
    o1_s0_out = tl.where(l1_pos, o1_s0 / l1, tl.zeros([BLOCK_DV], dtype=tl.float32))
    o1_s1_out = tl.where(l1_pos, o1_s1 / l1, tl.zeros([BLOCK_DV], dtype=tl.float32))
    o1_s2_out = tl.where(l1_pos, o1_s2 / l1, tl.zeros([BLOCK_DV], dtype=tl.float32))
    o1_s3_out = tl.where(l1_pos, o1_s3 / l1, tl.zeros([BLOCK_DV], dtype=tl.float32))
    tl.store(out_ptrs1 + d0 * stride_out_d, o1_s0_out.to(tl.bfloat16))
    tl.store(out_ptrs1 + d1 * stride_out_d, o1_s1_out.to(tl.bfloat16))
    tl.store(out_ptrs1 + d2 * stride_out_d, o1_s2_out.to(tl.bfloat16))
    tl.store(out_ptrs1 + d3 * stride_out_d, o1_s3_out.to(tl.bfloat16))

    # Head 2
    l2_pos = l2 > 0
    lse2 = tl.where(l2_pos, (m2 + tl.log(l2)) / ln2, neg_inf)
    lse_ptr2 = lse_ptr + pid_q * stride_lse_q + (heads_base + 2) * stride_lse_h
    tl.store(lse_ptr2, lse2)
    out_ptrs2 = out_ptr + pid_q * stride_out_q + (heads_base + 2) * stride_out_h
    o2_s0_out = tl.where(l2_pos, o2_s0 / l2, tl.zeros([BLOCK_DV], dtype=tl.float32))
    o2_s1_out = tl.where(l2_pos, o2_s1 / l2, tl.zeros([BLOCK_DV], dtype=tl.float32))
    o2_s2_out = tl.where(l2_pos, o2_s2 / l2, tl.zeros([BLOCK_DV], dtype=tl.float32))
    o2_s3_out = tl.where(l2_pos, o2_s3 / l2, tl.zeros([BLOCK_DV], dtype=tl.float32))
    tl.store(out_ptrs2 + d0 * stride_out_d, o2_s0_out.to(tl.bfloat16))
    tl.store(out_ptrs2 + d1 * stride_out_d, o2_s1_out.to(tl.bfloat16))
    tl.store(out_ptrs2 + d2 * stride_out_d, o2_s2_out.to(tl.bfloat16))
    tl.store(out_ptrs2 + d3 * stride_out_d, o2_s3_out.to(tl.bfloat16))

    # Head 3
    l3_pos = l3 > 0
    lse3 = tl.where(l3_pos, (m3 + tl.log(l3)) / ln2, neg_inf)
    lse_ptr3 = lse_ptr + pid_q * stride_lse_q + (heads_base + 3) * stride_lse_h
    tl.store(lse_ptr3, lse3)
    out_ptrs3 = out_ptr + pid_q * stride_out_q + (heads_base + 3) * stride_out_h
    o3_s0_out = tl.where(l3_pos, o3_s0 / l3, tl.zeros([BLOCK_DV], dtype=tl.float32))
    o3_s1_out = tl.where(l3_pos, o3_s1 / l3, tl.zeros([BLOCK_DV], dtype=tl.float32))
    o3_s2_out = tl.where(l3_pos, o3_s2 / l3, tl.zeros([BLOCK_DV], dtype=tl.float32))
    o3_s3_out = tl.where(l3_pos, o3_s3 / l3, tl.zeros([BLOCK_DV], dtype=tl.float32))
    tl.store(out_ptrs3 + d0 * stride_out_d, o3_s0_out.to(tl.bfloat16))
    tl.store(out_ptrs3 + d1 * stride_out_d, o3_s1_out.to(tl.bfloat16))
    tl.store(out_ptrs3 + d2 * stride_out_d, o3_s2_out.to(tl.bfloat16))
    tl.store(out_ptrs3 + d3 * stride_out_d, o3_s3_out.to(tl.bfloat16))


def _prepare_q_meta_from_indptr(qo_indptr: torch.Tensor, kv_indptr: torch.Tensor):
    # Build per-query arrays: kv_start[q], kv_max[q]
    qo_indptr_cpu = qo_indptr.to("cpu", non_blocking=False)
    kv_indptr_cpu = kv_indptr.to("cpu", non_blocking=False)
    len_indptr = qo_indptr_cpu.numel()
    total_q = int(qo_indptr_cpu[-1].item())
    q_kv_start = torch.empty(total_q, dtype=torch.int32)
    q_kv_max = torch.empty(total_q, dtype=torch.int32)
    for b in range(len_indptr - 1):
        q_start = int(qo_indptr_cpu[b].item())
        q_end = int(qo_indptr_cpu[b + 1].item())
        kv_start = int(kv_indptr_cpu[b].item())
        kv_end = int(kv_indptr_cpu[b + 1].item())
        q_len = q_end - q_start
        kv_len = kv_end - kv_start
        if q_len <= 0:
            continue
        delta = kv_len - q_len
        pos = torch.arange(q_len, dtype=torch.int32)
        kv_max = pos + 1 + int(delta)
        kv_max = torch.clamp(kv_max, min=0, max=kv_len)
        q_kv_start[q_start:q_end] = int(kv_start)
        q_kv_max[q_start:q_end] = kv_max
    return q_kv_start, q_kv_max


@torch.no_grad()
def run(q, k, v, qo_indptr, kv_indptr, sm_scale=None):
    if not torch.cuda.is_available():
        raise RuntimeError("CUDA is required to run Triton kernels. No CUDA device is available.")

    HEAD_DIM = 128
    NUM_QO = 32
    NUM_KV = 8
    RATIO = NUM_QO // NUM_KV  # 4

    inputs = [q, k, v, qo_indptr, kv_indptr]
    orig_devices = [t.device for t in inputs]
    target_device = None
    for t in inputs:
        if t.is_cuda:
            target_device = t.device
            break
    if target_device is None:
        target_device = torch.device("cuda")

    q = q.to(device=target_device, dtype=torch.bfloat16, non_blocking=True)
    k = k.to(device=target_device, dtype=torch.bfloat16, non_blocking=True)
    v = v.to(device=target_device, dtype=torch.bfloat16, non_blocking=True)
    qo_indptr = qo_indptr.to(device=target_device, dtype=torch.int32, non_blocking=True)
    kv_indptr = kv_indptr.to(device=target_device, dtype=torch.int32, non_blocking=True)

    total_q, num_qo_heads, head_dim = q.shape
    total_kv, num_kv_heads, _ = k.shape

    if num_qo_heads != NUM_QO:
        raise ValueError(f"num_qo_heads must be {NUM_QO}, got {num_qo_heads}")
    if num_kv_heads != NUM_KV:
        raise ValueError(f"num_kv_heads must be {NUM_KV}, got {num_kv_heads}")
    if head_dim != HEAD_DIM:
        raise ValueError(f"head_dim must be {HEAD_DIM}, got {head_dim}")

    if int(qo_indptr[-1].item()) != total_q:
        raise ValueError("Constraint violated: total_q must equal qo_indptr[-1]")
    if int(kv_indptr[-1].item()) != total_kv:
        raise ValueError("Constraint violated: total_kv must equal kv_indptr[-1]")

    if sm_scale is None:
        sm_scale = 1.0 / math.sqrt(HEAD_DIM)
    sm_scale = float(sm_scale)
    ln2 = float(math.log(2.0))

    # Prepare per-query kv_start and kv_max on CPU for simplicity, then move to target device
    qo_indptr_cpu = qo_indptr.to("cpu")
    kv_indptr_cpu = kv_indptr.to("cpu")
    q_kv_start_cpu, q_kv_max_cpu = _prepare_q_meta_from_indptr(qo_indptr_cpu, kv_indptr_cpu)
    q_kv_start = q_kv_start_cpu.to(device=target_device, non_blocking=True)
    q_kv_max = q_kv_max_cpu.to(device=target_device, non_blocking=True)

    out_gpu = torch.empty((total_q, NUM_QO, HEAD_DIM), dtype=torch.bfloat16, device=target_device)
    lse_gpu = torch.empty((total_q, NUM_QO), dtype=torch.float32, device=target_device)

    stride_q_q, stride_q_h, stride_q_d = q.stride()
    stride_k_k, stride_k_h, stride_k_d = k.stride()
    stride_v_k, stride_v_h, stride_v_d = v.stride()
    stride_out_q, stride_out_h, stride_out_d = out_gpu.stride()
    stride_lse_q, stride_lse_h = lse_gpu.stride()

    grid = (total_q, NUM_KV)
    BLOCK_N = 64
    BLOCK_DK = 32
    BLOCK_DV = 32

    gqa_ragged_prefill_causal_h32_kv8_d128_kernel[grid](
        q, k, v,
        stride_q_q, stride_q_h, stride_q_d,
        stride_k_k, stride_k_h, stride_k_d,
        stride_v_k, stride_v_h, stride_v_d,
        out_gpu, stride_out_q, stride_out_h, stride_out_d,
        lse_gpu, stride_lse_q, stride_lse_h,
        q_kv_start, q_kv_max,
        total_q,
        sm_scale, ln2,
        RATIO=RATIO, HEAD_DIM=HEAD_DIM,
        BLOCK_N=BLOCK_N, BLOCK_DK=BLOCK_DK, BLOCK_DV=BLOCK_DV,
        num_warps=4, num_stages=2,
    )

    out = out_gpu.to(orig_devices[0], non_blocking=True)
    lse = lse_gpu.to(orig_devices[0], non_blocking=True)
    return out, lse
scrolls · 386 lines total

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

Best evidence level for this revision: reported

JSON