Skip to content
KernelIndex
Search⌘K

gpt-o3 / tritondeaf62

gpt-o3_triton_deaf62 · gpt-o3 · triton · Apache-2.0

Use it

Vendorable · source mirrored · Apache-2.0View source →

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

main.py
curl "https://kernelindex.com/api/v1/implementations/flashinfer-gpt-o3-triton-deaf62?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:798c396221a454a79362500b188e1f32e6600949d92486a9d4ae8ac096c6f16e
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.

mmascores = tl.dot(q, tl.trans(k_tile)) * sm_scale # [BM, BN]
num-warps = 8NUM_WARPS = 8 # good default for B200 GPUs
online-softmaxm_new = tl.maximum(m_i, m_ij)
stages = 1num_stages=1,
tile-m = 64BLOCK_M = 64 # queries per block
tile-n = 64BLOCK_N = 64 # keys per block

Kernel source

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

# -----------------------------------------------------------------------------#
# Global compile-time constants                                                #
# -----------------------------------------------------------------------------#
NUM_QO_HEADS = 32
NUM_KV_HEADS = 4
HEAD_DIM     = 128
GQA_RATIO    = NUM_QO_HEADS // NUM_KV_HEADS

# Tunable tile sizes for B200
BLOCK_M   = 64     # queries  per block
BLOCK_N   = 64     # keys     per block
NUM_WARPS = 8      # good default for B200 GPUs


# -----------------------------------------------------------------------------#
# Triton kernel                                                                #
# -----------------------------------------------------------------------------#
@triton.jit
def _gqa_ragged_prefill_kernel(
    Q_ptr, K_ptr, V_ptr,                        # *bf16
    O_ptr, LSE_ptr,                             # *bf16 / *fp32
    q_start: tl.int32,                          # offset of first query token
    kv_start: tl.int32,                         # offset of first kv token
    q_len: tl.int32,                            # number of query tokens
    kv_len: tl.int32,                           # number of kv tokens
    delta: tl.int32,                            # kv_len - q_len
    sm_scale: tl.float32,                       # softmax scale
    inv_ln2: tl.float32,                        # 1 / ln(2)
    *,
    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,
):
    # ------------------------------------------------------------------#
    # Program IDs                                                       #
    # ------------------------------------------------------------------#
    pid_m = tl.program_id(0)     # query-block id
    pid_h = tl.program_id(1)     # qo-head id (0 … 31)

    # ------------------------------------------------------------------#
    # Index computations                                                #
    # ------------------------------------------------------------------#
    offs_m = pid_m * BLOCK_M + tl.arange(0, BLOCK_M)      # [BM]
    offs_n = tl.arange(0, BLOCK_N)                        # [BN]
    offs_d = tl.arange(0, HEAD_DIM)                       # [HD]

    row_mask = offs_m < q_len                             # [BM] bool

    qo_head = pid_h
    kv_head = qo_head // GQA_RATIO

    stride_q_token  = NUM_QO_HEADS * HEAD_DIM
    stride_kv_token = NUM_KV_HEADS * HEAD_DIM

    # ------------------------------------------------------------------#
    # Load Q                                                            #
    # ------------------------------------------------------------------#
    q_ptrs = (
        Q_ptr
        + (q_start + offs_m[:, None]) * stride_q_token
        + qo_head * HEAD_DIM
        + offs_d[None, :]
    )
    q = tl.load(q_ptrs, mask=row_mask[:, None], other=0).to(tl.float32)   # [BM, HD]

    # ------------------------------------------------------------------#
    # Online softmax initialisation                                     #
    # ------------------------------------------------------------------#
    NEG_INF = -1.0e30
    m_i = tl.full((BLOCK_M,), NEG_INF, dtype=tl.float32)
    l_i = tl.zeros((BLOCK_M,), dtype=tl.float32)
    acc = tl.zeros((BLOCK_M, HEAD_DIM), dtype=tl.float32)

    # ------------------------------------------------------------------#
    # Iterate over KV tiles                                             #
    # ------------------------------------------------------------------#
    kv_tile_start = tl.int32(0)
    while kv_tile_start < kv_len:
        k_ids = kv_tile_start + offs_n                         # [BN]
        k_valid = k_ids < kv_len                               # [BN] bool

        # ---- load K / V ---------------------------------------------
        k_ptrs = (
            K_ptr
            + (kv_start + k_ids[:, None]) * stride_kv_token
            + kv_head * HEAD_DIM
            + offs_d[None, :]
        )
        v_ptrs = (
            V_ptr
            + (kv_start + k_ids[:, None]) * stride_kv_token
            + kv_head * HEAD_DIM
            + offs_d[None, :]
        )
        k_tile = tl.load(k_ptrs, mask=k_valid[:, None], other=0).to(tl.float32)  # [BN, HD]
        v_tile = tl.load(v_ptrs, mask=k_valid[:, None], other=0).to(tl.float32)  # [BN, HD]

        # ---- attention scores ----------------------------------------
        scores = tl.dot(q, tl.trans(k_tile)) * sm_scale        # [BM, BN]

        # ---- causal masking ------------------------------------------
        allowed_k = offs_m + delta + 1                         # [BM]
        causal_mask = k_ids[None, :] >= allowed_k[:, None]     # [BM, BN]
        valid_mask = k_valid[None, :] & (~causal_mask) & row_mask[:, None]

        scores = tl.where(valid_mask, scores, NEG_INF)

        # ---- online softmax ------------------------------------------
        m_ij = tl.max(scores, axis=1)                          # [BM]
        m_new = tl.maximum(m_i, m_ij)

        exp_m_i    = tl.exp(m_i - m_new)
        exp_scores = tl.exp(scores - m_new[:, None]) * valid_mask.to(tl.float32)

        l_new = l_i * exp_m_i + tl.sum(exp_scores, axis=1)     # [BM]

        # update accumulator
        pv = tl.dot(exp_scores, v_tile)                        # [BM, HD]
        acc = (acc * (l_i * exp_m_i)[:, None] + pv) / l_new[:, None]

        m_i = m_new
        l_i = l_new

        kv_tile_start += BLOCK_N

    # ------------------------------------------------------------------#
    # Write back output & LSE                                           #
    # ------------------------------------------------------------------#
    o_ptrs = (
        O_ptr
        + (q_start + offs_m[:, None]) * stride_q_token
        + qo_head * HEAD_DIM
        + offs_d[None, :]
    )
    tl.store(o_ptrs, acc.to(tl.bfloat16), mask=row_mask[:, None])

    lse_vals = (m_i + tl.log(l_i)) * inv_ln2                 # [BM]
    lse_ptrs = (
        LSE_ptr
        + (q_start + offs_m) * NUM_QO_HEADS
        + qo_head
    )
    tl.store(lse_ptrs, lse_vals, mask=row_mask)


# -----------------------------------------------------------------------------#
# Python wrapper                                                               #
# -----------------------------------------------------------------------------#
@torch.no_grad()
def run(q, k, v, qo_indptr, kv_indptr, sm_scale=None):
    """
    Optimised Triton implementation of
    gqa_ragged_prefill_causal_h32_kv4_d128
    """
    # ---------------------------------------------------------------#
    # Device management                                              #
    # ---------------------------------------------------------------#
    if not torch.cuda.is_available():
        raise RuntimeError("CUDA device is required to run Triton kernels.")

    orig_device = q.device
    if orig_device.type == "cpu":
        q, k, v = q.cuda(), k.cuda(), v.cuda()
        qo_indptr, kv_indptr = qo_indptr.cuda(), kv_indptr.cuda()
    elif orig_device.type != "cuda":
        raise RuntimeError(f"Unsupported device type: {orig_device.type!r}")

    # ---------------------------------------------------------------#
    # Shape / constant checks                                        #
    # ---------------------------------------------------------------#
    total_q, num_qo_heads, head_dim = q.shape
    total_kv, num_kv_heads, _ = k.shape

    assert num_qo_heads == NUM_QO_HEADS, "num_qo_heads mismatch"
    assert num_kv_heads == NUM_KV_HEADS, "num_kv_heads mismatch"
    assert head_dim == HEAD_DIM, "head_dim mismatch"
    assert total_q == qo_indptr[-1].item(), "total_q != qo_indptr[-1]"
    assert total_kv == kv_indptr[-1].item(), "total_kv != kv_indptr[-1]"

    # ---------------------------------------------------------------#
    # Soft-max scale                                                 #
    # ---------------------------------------------------------------#
    if sm_scale is None:
        sm_scale = 1.0 / math.sqrt(head_dim)
    sm_scale = float(sm_scale)
    inv_ln2 = 1.0 / math.log(2.0)

    # ---------------------------------------------------------------#
    # Allocate outputs                                               #
    # ---------------------------------------------------------------#
    output = torch.empty(
        (total_q, NUM_QO_HEADS, HEAD_DIM),
        dtype=torch.bfloat16,
        device=q.device,
    )
    lse = torch.empty(
        (total_q, NUM_QO_HEADS),
        dtype=torch.float32,
        device=q.device,
    )

    # ---------------------------------------------------------------#
    # Launch kernel for each sequence                                #
    # ---------------------------------------------------------------#
    batch_size = qo_indptr.numel() - 1
    for b in range(batch_size):
        q_start  = int(qo_indptr[b].item())
        q_end    = int(qo_indptr[b + 1].item())
        kv_start = int(kv_indptr[b].item())
        kv_end   = int(kv_indptr[b + 1].item())

        if q_start >= q_end or kv_start >= kv_end:
            continue  # empty slice

        q_len  = q_end  - q_start
        kv_len = kv_end - kv_start
        delta  = kv_len - q_len

        grid_m = triton.cdiv(q_len, BLOCK_M)
        grid = (grid_m, NUM_QO_HEADS)

        _gqa_ragged_prefill_kernel[grid](
            q, k, v,
            output, lse,
            q_start, kv_start,
            q_len, kv_len,
            delta,
            sm_scale,
            inv_ln2,
            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=1,
        )

    # ---------------------------------------------------------------#
    # Move outputs back to original device                           #
    # ---------------------------------------------------------------#
    if orig_device.type == "cpu":
        output = output.cpu()
        lse = lse.cpu()

    return output, lse
scrolls · 255 lines total

Source code from the importing source · Apache-2.0

No published measurement for this revision

JSON