Skip to content
KernelIndex
Search⌘K

gpt-5 / triton41ae45

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

Use it

Vendorable · source mirrored · Apache-2.0View source →

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

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

Benchmark evidence

No published measurement for this revision.

No evidence · How evidence levels are derived →

Source and license

sourcehttps://huggingface.co/datasets/flashinfer-ai/flashinfer-trace
commitda915083d4c7c5e61aa3005e3d17ae488e0fc71c
revision digestsha256:dceb94f9dc4230e0173844b6e4def21143d27131ec2fb6499eb03ee2e21e1200
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.

mmaqk = tl.dot(q_tile, tl.trans(k_tile))
num-warps = 8num_warps = 8
stages = 2num_stages = 2
tile-m = 32BLOCK_M = 32
tile-n = 128BLOCK_N = 128

Kernel source

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


@triton.jit
def gqa_ragged_prefill_causal_h32_kv4_d128_kernel(
    Q_ptr, K_ptr, V_ptr, O_ptr, LSE_ptr,
    qo_indptr_ptr, kv_indptr_ptr,
    total_q, total_kv,
    sm_scale,
    stride_q0, stride_q1, stride_q2,
    stride_k0, stride_k1, stride_k2,
    stride_v0, stride_v1, stride_v2,
    stride_o0, stride_o1, stride_o2,
    stride_lse0, stride_lse1,
    BLOCK_M: tl.constexpr, BLOCK_N: tl.constexpr,
    HEAD_DIM: tl.constexpr, NUM_QO_HEADS: tl.constexpr, NUM_KV_HEADS: tl.constexpr, GQA_RATIO: tl.constexpr
):
    pid_seq = tl.program_id(0)    # sequence id
    pid_kvh = tl.program_id(1)    # kv head id
    pid_mblk = tl.program_id(2)   # query block id within sequence

    # Load sequence boundaries
    q_start = tl.load(qo_indptr_ptr + pid_seq)
    q_end = tl.load(qo_indptr_ptr + pid_seq + 1)
    kv_start = tl.load(kv_indptr_ptr + pid_seq)
    kv_end = tl.load(kv_indptr_ptr + pid_seq + 1)

    q_len = q_end - q_start
    kv_len = kv_end - kv_start

    # Offsets within sequence for queries
    m_offsets = pid_mblk * BLOCK_M + tl.arange(0, BLOCK_M)
    m_mask = m_offsets < q_len
    q_abs = q_start + m_offsets

    # Causal delta: kv_len - q_len
    delta = kv_len - q_len

    # Per-row kv cap: n_cap = m + 1 + delta
    n_cap = m_offsets + (1 + delta)
    has_attn_row = n_cap > 0

    # Dimension offsets
    d_offsets = tl.arange(0, HEAD_DIM)

    kv_h = pid_kvh
    qo_h_base = pid_kvh * GQA_RATIO

    NEG_INF = float("-inf")
    INV_LN2 = 1.4426950408889634  # 1 / ln(2)
    sm_scale_f32 = tl.full([1], sm_scale, tl.float32)[0]

    # Iterate over the 8 Qo-heads mapped to this kv head
    for h in tl.static_range(GQA_RATIO):
        qo_h = qo_h_base + h

        # Load Q tile [M, D] in f32
        q_ptrs = Q_ptr + q_abs[:, None] * stride_q0 + qo_h * stride_q1 + d_offsets[None, :] * stride_q2
        q_tile = tl.load(q_ptrs, mask=m_mask[:, None], other=0.0).to(tl.float32)

        # Pass 1: compute per-row max (m_i) over all K tiles with causal mask
        m_i = tl.full([BLOCK_M], NEG_INF, tl.float32)

        n_start = 0
        while n_start < kv_len:
            n_offsets = n_start + tl.arange(0, BLOCK_N)
            n_inbounds = n_offsets < kv_len

            # Load K tile [N, D] for this kv head
            k_ptrs = K_ptr + (kv_start + n_offsets)[:, None] * stride_k0 + kv_h * stride_k1 + d_offsets[None, :] * stride_k2
            k_tile = tl.load(k_ptrs, mask=n_inbounds[:, None], other=0.0).to(tl.float32)

            # QK^T
            qk = tl.dot(q_tile, tl.trans(k_tile))
            qk_scaled = qk * sm_scale_f32

            # Causal mask for this tile
            n_base = n_offsets[None, :]            # [1, N]
            n_cap_broadcast = n_cap[:, None]       # [M, 1]
            causal_mask = (n_base < n_cap_broadcast) & n_inbounds[None, :] & m_mask[:, None]

            # Compute tile max with mask
            qk_masked = tl.where(causal_mask, qk_scaled, NEG_INF)
            tile_max = tl.max(qk_masked, axis=1)
            m_i = tl.maximum(m_i, tile_max)

            n_start += BLOCK_N

        # Pass 2: compute sum of exp and weighted value accumulation
        l_i = tl.zeros([BLOCK_M], tl.float32)
        acc = tl.zeros([BLOCK_M, HEAD_DIM], tl.float32)

        n_start = 0
        while n_start < kv_len:
            n_offsets = n_start + tl.arange(0, BLOCK_N)
            n_inbounds = n_offsets < kv_len

            # Load K and V tiles for this kv head [N, D]
            k_ptrs = K_ptr + (kv_start + n_offsets)[:, None] * stride_k0 + kv_h * stride_k1 + d_offsets[None, :] * stride_k2
            v_ptrs = V_ptr + (kv_start + n_offsets)[:, None] * stride_v0 + kv_h * stride_v1 + d_offsets[None, :] * stride_v2

            k_tile = tl.load(k_ptrs, mask=n_inbounds[:, None], other=0.0).to(tl.float32)
            v_tile = tl.load(v_ptrs, mask=n_inbounds[:, None], other=0.0).to(tl.float32)

            # QK^T
            qk = tl.dot(q_tile, tl.trans(k_tile))
            qk_scaled = qk * sm_scale_f32

            # Causal mask
            n_base = n_offsets[None, :]            # [1, N]
            n_cap_broadcast = n_cap[:, None]       # [M, 1]
            causal_mask = (n_base < n_cap_broadcast) & n_inbounds[None, :] & m_mask[:, None]

            # Stable logits with global row max m_i
            stable_logits = qk_scaled - m_i[:, None]
            stable_logits = tl.where(causal_mask, stable_logits, NEG_INF)

            # Probabilities and accumulation
            p = tl.exp(stable_logits)
            l_i += tl.sum(p, axis=1)
            acc += tl.dot(p, v_tile)

            n_start += BLOCK_N

        # Build store mask: only rows with queries and at least one valid key
        m_store_mask = m_mask & has_attn_row

        # Normalize output
        l_i_safe = tl.where(m_store_mask, l_i, 1.0)
        out_tile = acc / l_i_safe[:, None]

        # Store output
        o_ptrs = O_ptr + q_abs[:, None] * stride_o0 + qo_h * stride_o1 + d_offsets[None, :] * stride_o2
        tl.store(o_ptrs, out_tile.to(tl.bfloat16), mask=m_store_mask[:, None])

        # LSE base-2: (log(sum(exp)) + m_i) / ln(2)
        lse_vals = (tl.log(l_i) + m_i) * INV_LN2
        lse_ptrs = LSE_ptr + q_abs * stride_lse0 + qo_h * stride_lse1
        tl.store(lse_ptrs, lse_vals, mask=m_store_mask)


def _ceil_div_int(a: int, b: int) -> int:
    return (a + b - 1) // b


def run(q, k, v, qo_indptr, kv_indptr, sm_scale=None):
    # Validate CUDA availability
    cuda_available = torch.cuda.is_available()
    devices = {
        "q": q.device,
        "k": k.device,
        "v": v.device,
        "qo_indptr": qo_indptr.device,
        "kv_indptr": kv_indptr.device,
    }
    target_device = devices["q"]
    if not cuda_available:
        if any(t.is_cuda for t in [q, k, v, qo_indptr, kv_indptr]):
            raise RuntimeError("CUDA is not available but GPU tensors were provided.")
        raise RuntimeError("CUDA is required to run Triton kernels.")

    # Shapes and checks
    total_q, num_qo_heads, head_dim = q.shape
    total_kv, num_kv_heads, _ = k.shape
    len_indptr = qo_indptr.shape[0]

    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, "head_dim must be 128"
    assert total_q == int(qo_indptr[-1].item()), "total_q must equal qo_indptr[-1]"
    assert total_kv == int(kv_indptr[-1].item()), "total_kv must equal kv_indptr[-1]"
    assert k.shape == v.shape, "k and v must have same shape"
    assert qo_indptr.shape[0] == kv_indptr.shape[0], "qo_indptr and kv_indptr must have same length"

    # Default sm_scale
    if sm_scale is None:
        sm_scale = 1.0 / math.sqrt(head_dim)
    # Cast to float32 to avoid 64-bit scalar promotion differences
    sm_scale = float(torch.tensor(sm_scale, dtype=torch.float32).item())

    # Dtype checks
    if q.dtype != torch.bfloat16 or k.dtype != torch.bfloat16 or v.dtype != torch.bfloat16:
        raise TypeError("q, k, v must be torch.bfloat16")
    if qo_indptr.dtype != torch.int32 or kv_indptr.dtype != torch.int32:
        raise TypeError("qo_indptr and kv_indptr must be torch.int32")

    compute_device = torch.device("cuda")

    # Move to CUDA
    q_dev = q if q.device.type == "cuda" else q.to(compute_device, non_blocking=True)
    k_dev = k if k.device.type == "cuda" else k.to(compute_device, non_blocking=True)
    v_dev = v if v.device.type == "cuda" else v.to(compute_device, non_blocking=True)
    qo_indptr_dev = qo_indptr if qo_indptr.device.type == "cuda" else qo_indptr.to(compute_device, non_blocking=True)
    kv_indptr_dev = kv_indptr if kv_indptr.device.type == "cuda" else kv_indptr.to(compute_device, non_blocking=True)

    # Prepare outputs on device
    out_dev = torch.zeros((total_q, num_qo_heads, head_dim), dtype=torch.bfloat16, device=compute_device)
    lse_dev = torch.full((total_q, num_qo_heads), float("-inf"), dtype=torch.float32, device=compute_device)

    # Early exit if no sequences
    num_seqs = len_indptr - 1
    if num_seqs <= 0 or total_q == 0 or total_kv == 0:
        target_out = out_dev if target_device.type == "cuda" else out_dev.to(target_device, non_blocking=True)
        target_lse = lse_dev if target_device.type == "cuda" else lse_dev.to(target_device, non_blocking=True)
        return target_out, target_lse

    # Constants
    GQA_RATIO = 8
    NUM_QO_HEADS = 32
    NUM_KV_HEADS = 4
    HEAD_DIM = 128

    # Block sizes tuned conservatively for B200
    BLOCK_M = 32
    BLOCK_N = 128

    # Number of M blocks per sequence, use max across sequences for grid; masking handles others
    qo_indptr_cpu = qo_indptr_dev.detach().cpu()
    q_lengths = (qo_indptr_cpu[1:] - qo_indptr_cpu[:-1]).to(torch.int64)
    if q_lengths.numel() > 0:
        max_q_blocks = int(((q_lengths + (BLOCK_M - 1)) // BLOCK_M).max().item())
        if max_q_blocks <= 0:
            max_q_blocks = 1
    else:
        max_q_blocks = 1

    # Strides
    stride_q0, stride_q1, stride_q2 = q_dev.stride()
    stride_k0, stride_k1, stride_k2 = k_dev.stride()
    stride_v0, stride_v1, stride_v2 = v_dev.stride()
    stride_o0, stride_o1, stride_o2 = out_dev.stride()
    stride_lse0, stride_lse1 = lse_dev.stride()

    grid = (num_seqs, NUM_KV_HEADS, max_q_blocks)
    num_warps = 8
    num_stages = 2

    gqa_ragged_prefill_causal_h32_kv4_d128_kernel[grid](
        q_dev, k_dev, v_dev, out_dev, lse_dev,
        qo_indptr_dev, kv_indptr_dev,
        total_q, total_kv,
        sm_scale,
        stride_q0, stride_q1, stride_q2,
        stride_k0, stride_k1, stride_k2,
        stride_v0, stride_v1, stride_v2,
        stride_o0, stride_o1, stride_o2,
        stride_lse0, stride_lse1,
        BLOCK_M=BLOCK_M, BLOCK_N=BLOCK_N,
        HEAD_DIM=HEAD_DIM, NUM_QO_HEADS=NUM_QO_HEADS, NUM_KV_HEADS=NUM_KV_HEADS, GQA_RATIO=GQA_RATIO,
        num_warps=num_warps, num_stages=num_stages
    )

    # Move outputs back to original device of q
    if target_device.type != "cuda":
        out_host = out_dev.to(target_device, non_blocking=True)
        lse_host = lse_dev.to(target_device, non_blocking=True)
    else:
        out_host = out_dev
        lse_host = lse_dev

    return out_host, lse_host
scrolls · 264 lines total

Source code from the importing source · Apache-2.0

No published measurement for this revision

JSON