kernelbench.com

KernelBench hard · H100

KDA CUTLASS GPT-5.6 Sol

1.32%geomean peak fraction across shapes

manually audited: clean

harnesscodexagent session27mtotal wall28mcheck42sbenchmark7soutput tokens54,714gpu-lock wait3sgpu-lock held7mregimecompute

Per-shape vs governing ceilingeach shape graded against whichever binds — bf16 compute or HBM bandwidth

2×1024×8×128×128×640.179 ms1.6%0.14 TB/s · 7% of 2.0 TB/s HBM · also 12 TFLOPS (2% of compute)
2×2048×8×128×128×640.318 ms1.8%0.16 TB/s · 8% of 2.0 TB/s HBM · also 13 TFLOPS (2% of compute)
1×4096×8×128×128×640.432 ms1.3%0.12 TB/s · 6% of 2.0 TB/s HBM · also 10 TFLOPS (1% of compute)
1×2048×4×128×128×640.176 ms0.8%0.07 TB/s · 4% of 2.0 TB/s HBM · also 6 TFLOPS (1% of compute)

compute-bound memory-bound · bar + right column = official fraction of the ceiling (the geomean input)

geomean(1.6% · 1.8% · 1.3% · 0.8%) = 1.3%

Kernel source (redacted)
"""SM90-oriented Triton forward kernels for Kimi Delta Attention.

Each 64-token chunk is prepared independently with tensor-core products and
a four-block triangular inverse.  A final persistent kernel owns one
(batch, head, value-tile), fuses the inter-chunk recurrence with the output
projection, and keeps its 128 x BV state on chip for the complete sequence.
"""
from __future__ import annotations

import torch
import torch.nn as nn
import triton
import triton.language as tl

PAIR_WARPS = 4
PAIR_STAGES = 3
SOLVE_WARPS = 2
SOLVE_STAGES = 1
PROJECT_WARPS = 4
PROJECT_STAGES = 3
STATE_BV = 16
STATE_WARPS = 4
OUTPUT_BV = 128
OUTPUT_WARPS = 4
OUTPUT_STAGES = 3


@triton.jit
def _delta_forward_kernel(
    q_ptr,
    k_ptr,
    v_ptr,
    g_ptr,
    beta_ptr,
    out_ptr,
    T,
    H: tl.constexpr,
    D: tl.constexpr,
    BV: tl.constexpr,
    SCALE: tl.constexpr,
):
    pid_v = tl.program_id(0)
    pid_bh = tl.program_id(1)
    b = pid_bh // H
    h = pid_bh - b * H

    ik = tl.arange(0, D)
    iv = pid_v * BV + tl.arange(0, BV)
    mv = iv < D

    # This tensor is distributed over the participating warps and stays in
    # registers for the dynamic sequence loop.
    state = tl.zeros((D, BV), tl.float32)

    t = 0
    while t < T:
        qk_base = ((b * T + t) * H + h) * D
        v_base = qk_base
        bh_base = (b * T + t) * H + h

        q = tl.load(q_ptr + qk_base + ik).to(tl.float32)
        k = tl.load(k_ptr + qk_base + ik).to(tl.float32)
        decay = tl.exp(tl.load(g_ptr + qk_base + ik))
        value = tl.load(v_ptr + v_base + iv, mask=mv, other=0.0).to(tl.float32)
        beta = tl.load(beta_ptr + bh_base).to(tl.float32)

        state *= decay[:, None]
        prediction = tl.sum(state * k[:, None], axis=0)
        residual = beta * (value - prediction)
        state += k[:, None] * residual[None, :]
        output = tl.sum(state * q[:, None], axis=0) * SCALE

        tl.store(out_ptr + v_base + iv, output, mask=mv)
        t += 1


@triton.jit
def _local_prefix_kernel(
    src,
    dst,
    T,
    H: tl.constexpr,
    D: tl.constexpr,
    BT: tl.constexpr,
    BS: tl.constexpr,
):
    pid_s = tl.program_id(0)
    pid_t = tl.program_id(1)
    pid_bh = tl.program_id(2)
    b = pid_bh // H
    h = pid_bh - b * H
    base = (b * T * H + h) * D
    p_src = tl.make_block_ptr(
        src + base, (T, D), (H * D, 1),
        (pid_t * BT, pid_s * BS), (BT, BS), (1, 0),
    )
    p_dst = tl.make_block_ptr(
        dst + base, (T, D), (H * D, 1),
        (pid_t * BT, pid_s * BS), (BT, BS), (1, 0),
    )
    x = tl.load(p_src)
    tl.store(p_dst, tl.cumsum(x, axis=0))


@triton.jit
def _build_chunks_kernel(
    q_ptr,
    k_ptr,
    v_ptr,
    gc_ptr,
    beta_ptr,
    score_ptr,
    w_ptr,
    u_ptr,
    kg_ptr,
    T,
    H: tl.constexpr,
    D: tl.constexpr,
    BT: tl.constexpr,
    SCALE: tl.constexpr,
):
    pid_t = tl.program_id(0)
    pid_bh = tl.program_id(1)
    b = pid_bh // H
    h = pid_bh - b * H
    tc = pid_t * BT

    qk_base = (b * T * H + h) * D
    s_base = (b * T * H + h) * BT
    beta_base = b * T * H + h

    p_q = tl.make_block_ptr(q_ptr + qk_base, (T, D), (H * D, 1), (tc, 0), (BT, D), (1, 0))
    p_k = tl.make_block_ptr(k_ptr + qk_base, (T, D), (H * D, 1), (tc, 0), (BT, D), (1, 0))
    p_v = tl.make_block_ptr(v_ptr + qk_base, (T, D), (H * D, 1), (tc, 0), (BT, D), (1, 0))
    p_g = tl.make_block_ptr(gc_ptr + qk_base, (T, D), (H * D, 1), (tc, 0), (BT, D), (1, 0))
    p_beta = tl.make_block_ptr(beta_ptr + beta_base, (T,), (H,), (tc,), (BT,), (0,))

    q = tl.load(p_q)
    k = tl.load(p_k)
    g = tl.load(p_g)
    beta = tl.load(p_beta).to(tl.float32)

    # Centering the exponentials does not change either product and keeps both
    # factors well-scaled even for an unusually lopsided chunk.
    anchor = tl.sum(tl.where(tl.arange(0, BT)[:, None] == BT // 2, g, 0.0), axis=0)
    ep = tl.exp(g - anchor[None, :])
    em = tl.exp(anchor[None, :] - g)
    qg = q.to(tl.float32) * ep
    kg_left = k.to(tl.float32) * ep
    kg_right = k.to(tl.float32) * em

    scores = tl.dot(qg, tl.trans(kg_right), input_precision="tf32") * SCALE
    interaction = tl.dot(kg_left, tl.trans(kg_right), input_precision="tf32")

    ii = tl.arange(0, BT)
    lower = ii[:, None] > ii[None, :]
    causal = ii[:, None] >= ii[None, :]
    interaction = tl.where(lower, interaction * beta[:, None], 0.0)
    scores = tl.where(causal, scores, 0.0)

    # Invert the unit lower-triangular system by forward substitution.  Each
    # completed row is immediately visible to subsequent rows in `inv`.
    inv = -interaction
    for row_idx in range(2, BT):
        row = -tl.sum(
            tl.where(ii[:, None] == row_idx, interaction, 0.0), axis=0
        )
        row += tl.sum(row[:, None] * inv, axis=0)
        inv = tl.where((ii == row_idx)[:, None], row[None, :], inv)
    inv += (ii[:, None] == ii[None, :])
    transform = inv * beta[None, :]

    p_score = tl.make_block_ptr(
        score_ptr + s_base, (T, BT), (H * BT, 1), (tc, 0), (BT, BT), (1, 0)
    )
    tl.store(p_score, scores.to(tl.bfloat16))

    # These three intermediates are deliberately bf16: their consumers are
    # tensor-core products and the fp32 state holds the inter-chunk accuracy.
    v = tl.load(p_v)
    u = tl.dot(transform.to(tl.bfloat16), v)
    p_u = tl.make_block_ptr(u_ptr + qk_base, (T, D), (H * D, 1), (tc, 0), (BT, D), (1, 0))
    tl.store(p_u, u.to(tl.bfloat16))

    k_exp = (k.to(tl.float32) * tl.exp(g)).to(tl.bfloat16)
    w = tl.dot(transform.to(tl.bfloat16), k_exp)
    p_w = tl.make_block_ptr(w_ptr + qk_base, (T, D), (H * D, 1), (tc, 0), (BT, D), (1, 0))
    tl.store(p_w, w.to(tl.bfloat16))

    g_last = tl.sum(tl.where(ii[:, None] == BT - 1, g, 0.0), axis=0)
    k_tail = k.to(tl.float32) * tl.exp(g_last[None, :] - g)
    p_kg = tl.make_block_ptr(kg_ptr + qk_base, (T, D), (H * D, 1), (tc, 0), (BT, D), (1, 0))
    tl.store(p_kg, k_tail.to(tl.bfloat16))


@triton.jit
def _pair_matrices_kernel(
    q_ptr,
    k_ptr,
    g_ptr,
    gc_ptr,
    beta_ptr,
    score_ptr,
    lower_ptr,
    T,
    H: tl.constexpr,
    D: tl.constexpr,
    BT: tl.constexpr,
    SCALE: tl.constexpr,
):
    pid_t = tl.program_id(0)
    pid_bh = tl.program_id(1)
    b = pid_bh // H
    h = pid_bh - b * H
    tc = pid_t * BT
    qk_base = (b * T * H + h) * D
    m_base = (b * T * H + h) * BT
    beta_base = b * T * H + h

    p_q = tl.make_block_ptr(q_ptr + qk_base, (T, D), (H * D, 1), (tc, 0), (BT, D), (1, 0))
    p_k = tl.make_block_ptr(k_ptr + qk_base, (T, D), (H * D, 1), (tc, 0), (BT, D), (1, 0))
    p_g = tl.make_block_ptr(g_ptr + qk_base, (T, D), (H * D, 1), (tc, 0), (BT, D), (1, 0))
    p_gc = tl.make_block_ptr(gc_ptr + qk_base, (T, D), (H * D, 1), (tc, 0), (BT, D), (1, 0))
    p_b = tl.make_block_ptr(beta_ptr + beta_base, (T,), (H,), (tc,), (BT,), (0,))
    q = tl.load(p_q)
    k = tl.load(p_k)
    g = tl.cumsum(tl.load(p_g), axis=0)
    tl.store(p_gc, g)
    beta = tl.load(p_b).to(tl.float32)

    rows = tl.arange(0, BT)
    anchor = tl.sum(tl.where(rows[:, None] == BT // 2, g, 0.0), axis=0)
    ep = tl.exp(g - anchor[None, :])
    em = tl.exp(anchor[None, :] - g)
    right = k.to(tl.float32) * em
    score = tl.dot(q.to(tl.float32) * ep, tl.trans(right), input_precision="tf32") * SCALE
    lower = tl.dot(k.to(tl.float32) * ep, tl.trans(right), input_precision="tf32")
    score = tl.where(rows[:, None] >= rows[None, :], score, 0.0)
    lower = tl.where(rows[:, None] > rows[None, :], lower * beta[:, None], 0.0)

    p_score = tl.make_block_ptr(score_ptr + m_base, (T, BT), (H * BT, 1), (tc, 0), (BT, BT), (1, 0))
    p_lower = tl.make_block_ptr(lower_ptr + m_base, (T, BT), (H * BT, 1), (tc, 0), (BT, BT), (1, 0))
    tl.store(p_score, score.to(tl.bfloat16))
    tl.store(p_lower, lower)


@triton.jit
def _load_matrix16(ptr, base, T, H: tl.constexpr, BT: tl.constexpr,
                   BC: tl.constexpr, tc, RO: tl.constexpr, CO: tl.constexpr):
    p = tl.make_block_ptr(
        ptr + base, (T, BT), (H * BT, 1),
        (tc + RO * BC, CO * BC), (BC, BC), (1, 0),
    )
    return tl.load(p)


@triton.jit
def _load_beta16(ptr, base, T, H: tl.constexpr, BC: tl.constexpr,
                 tc, BLOCK: tl.constexpr):
    p = tl.make_block_ptr(ptr + base, (T,), (H,), (tc + BLOCK * BC,), (BC,), (0,))
    return tl.load(p).to(tl.float32)


@triton.jit
def _store_matrix16(ptr, x, base, T, H: tl.constexpr, BT: tl.constexpr,
                    BC: tl.constexpr, tc, RO: tl.constexpr, CO: tl.constexpr):
    p = tl.make_block_ptr(
        ptr + base, (T, BT), (H * BT, 1),
        (tc + RO * BC, CO * BC), (BC, BC), (1, 0),
    )
    tl.store(p, x.to(tl.bfloat16))


@triton.jit
def _solve64_kernel(
    lower_ptr,
    beta_ptr,
    transform_ptr,
    T,
    H: tl.constexpr,
    BT: tl.constexpr,
    BC: tl.constexpr,
):
    pid_t = tl.program_id(0)
    pid_bh = tl.program_id(1)
    b = pid_bh // H
    h = pid_bh - b * H
    tc = pid_t * BT
    m_base = (b * T * H + h) * BT
    beta_base = b * T * H + h

    l00 = _load_matrix16(lower_ptr, m_base, T, H, BT, BC, tc, 0, 0)
    l10 = _load_matrix16(lower_ptr, m_base, T, H, BT, BC, tc, 1, 0)
    l11 = _load_matrix16(lower_ptr, m_base, T, H, BT, BC, tc, 1, 1)
    l20 = _load_matrix16(lower_ptr, m_base, T, H, BT, BC, tc, 2, 0)
    l21 = _load_matrix16(lower_ptr, m_base, T, H, BT, BC, tc, 2, 1)
    l22 = _load_matrix16(lower_ptr, m_base, T, H, BT, BC, tc, 2, 2)
    l30 = _load_matrix16(lower_ptr, m_base, T, H, BT, BC, tc, 3, 0)
    l31 = _load_matrix16(lower_ptr, m_base, T, H, BT, BC, tc, 3, 1)
    l32 = _load_matrix16(lower_ptr, m_base, T, H, BT, BC, tc, 3, 2)
    l33 = _load_matrix16(lower_ptr, m_base, T, H, BT, BC, tc, 3, 3)

    ii = tl.arange(0, BC)
    strict = ii[:, None] > ii[None, :]
    ident = ii[:, None] == ii[None, :]
    a00 = -tl.where(strict, l00, 0.0)
    a11 = -tl.where(strict, l11, 0.0)
    a22 = -tl.where(strict, l22, 0.0)
    a33 = -tl.where(strict, l33, 0.0)

    for row_idx in range(2, BC):
        r0 = -tl.sum(tl.where(ii[:, None] == row_idx, l00, 0.0), axis=0)
        r1 = -tl.sum(tl.where(ii[:, None] == row_idx, l11, 0.0), axis=0)
        r2 = -tl.sum(tl.where(ii[:, None] == row_idx, l22, 0.0), axis=0)
        r3 = -tl.sum(tl.where(ii[:, None] == row_idx, l33, 0.0), axis=0)
        r0 += tl.sum(r0[:, None] * a00, axis=0)
        r1 += tl.sum(r1[:, None] * a11, axis=0)
        r2 += tl.sum(r2[:, None] * a22, axis=0)
        r3 += tl.sum(r3[:, None] * a33, axis=0)
        select = (ii == row_idx)[:, None]
        a00 = tl.where(select, r0[None, :], a00)
        a11 = tl.where(select, r1[None, :], a11)
        a22 = tl.where(select, r2[None, :], a22)
        a33 = tl.where(select, r3[None, :], a33)
    a00 += ident
    a11 += ident
    a22 += ident
    a33 += ident

    a10 = -tl.dot(tl.dot(a11, l10, input_precision="tf32"), a00, input_precision="tf32")
    a21 = -tl.dot(tl.dot(a22, l21, input_precision="tf32"), a11, input_precision="tf32")
    a20 = -tl.dot(
        a22,
        tl.dot(l20, a00, input_precision="tf32") + tl.dot(l21, a10, input_precision="tf32"),
        input_precision="tf32",
    )
    a32 = -tl.dot(tl.dot(a33, l32, input_precision="tf32"), a22, input_precision="tf32")
    a31 = -tl.dot(
        a33,
        tl.dot(l31, a11, input_precision="tf32") + tl.dot(l32, a21, input_precision="tf32"),
        input_precision="tf32",
    )
    a30 = -tl.dot(
        a33,
        tl.dot(l30, a00, input_precision="tf32")
        + tl.dot(l31, a10, input_precision="tf32")
        + tl.dot(l32, a20, input_precision="tf32"),
        input_precision="tf32",
    )

    b0 = _load_beta16(beta_ptr, beta_base, T, H, BC, tc, 0)
    b1 = _load_beta16(beta_ptr, beta_base, T, H, BC, tc, 1)
    b2 = _load_beta16(beta_ptr, beta_base, T, H, BC, tc, 2)
    b3 = _load_beta16(beta_ptr, beta_base, T, H, BC, tc, 3)

    _store_matrix16(transform_ptr, a00 * b0[None, :], m_base, T, H, BT, BC, tc, 0, 0)
    _store_matrix16(transform_ptr, a10 * b0[None, :], m_base, T, H, BT, BC, tc, 1, 0)
    _store_matrix16(transform_ptr, a11 * b1[None, :], m_base, T, H, BT, BC, tc, 1, 1)
    _store_matrix16(transform_ptr, a20 * b0[None, :], m_base, T, H, BT, BC, tc, 2, 0)
    _store_matrix16(transform_ptr, a21 * b1[None, :], m_base, T, H, BT, BC, tc, 2, 1)
    _store_matrix16(transform_ptr, a22 * b2[None, :], m_base, T, H, BT, BC, tc, 2, 2)
    _store_matrix16(transform_ptr, a30 * b0[None, :], m_base, T, H, BT, BC, tc, 3, 0)
    _store_matrix16(transform_ptr, a31 * b1[None, :], m_base, T, H, BT, BC, tc, 3, 1)
    _store_matrix16(transform_ptr, a32 * b2[None, :], m_base, T, H, BT, BC, tc, 3, 2)
    _store_matrix16(transform_ptr, a33 * b3[None, :], m_base, T, H, BT, BC, tc, 3, 3)


@triton.jit
def _project_chunks_kernel(
    k_ptr,
    v_ptr,
    gc_ptr,
    transform_ptr,
    w_ptr,
    u_ptr,
    kg_ptr,
    T,
    H: tl.constexpr,
    D: tl.constexpr,
    BT: tl.constexpr,
):
    pid_t = tl.program_id(0)
    pid_bh = tl.program_id(1)
    b = pid_bh // H
    h = pid_bh - b * H
    tc = pid_t * BT
    qk_base = (b * T * H + h) * D
    m_base = (b * T * H + h) * BT
    p_a = tl.make_block_ptr(transform_ptr + m_base, (T, BT), (H * BT, 1), (tc, 0), (BT, BT), (1, 0))
    a = tl.load(p_a)
    ii = tl.arange(0, BT)
    a = tl.where(ii[:, None] >= ii[None, :], a, 0.0)

    for block in range(2):
        d0 = block * 64
        p_k = tl.make_block_ptr(k_ptr + qk_base, (T, D), (H * D, 1), (tc, d0), (BT, 64), (1, 0))
        p_v = tl.make_block_ptr(v_ptr + qk_base, (T, D), (H * D, 1), (tc, d0), (BT, 64), (1, 0))
        p_g = tl.make_block_ptr(gc_ptr + qk_base, (T, D), (H * D, 1), (tc, d0), (BT, 64), (1, 0))
        k = tl.load(p_k)
        v = tl.load(p_v)
        g = tl.load(p_g)
        u = tl.dot(a, v)
        w = tl.dot(a, (k.to(tl.float32) * tl.exp(g)).to(tl.bfloat16))
        last = tl.sum(tl.where(ii[:, None] == BT - 1, g, 0.0), axis=0)
        kg = k.to(tl.float32) * tl.exp(last[None, :] - g)

        p_u = tl.make_block_ptr(u_ptr + qk_base, (T, D), (H * D, 1), (tc, d0), (BT, 64), (1, 0))
        p_w = tl.make_block_ptr(w_ptr + qk_base, (T, D), (H * D, 1), (tc, d0), (BT, 64), (1, 0))
        p_kg = tl.make_block_ptr(kg_ptr + qk_base, (T, D), (H * D, 1), (tc, d0), (BT, 64), (1, 0))
        tl.store(p_u, u.to(tl.bfloat16))
        tl.store(p_w, w.to(tl.bfloat16))
        tl.store(p_kg, kg.to(tl.bfloat16))


@triton.jit
def _chunks_recurrence_kernel(
    q_ptr,
    gc_ptr,
    score_ptr,
    w_ptr,
    u_ptr,
    kg_ptr,
    out_ptr,
    T,
    H: tl.constexpr,
    D: tl.constexpr,
    BT: tl.constexpr,
    BV: tl.constexpr,
    SCALE: tl.constexpr,
):
    pid_v = tl.program_id(0)
    pid_bh = tl.program_id(1)
    b = pid_bh // H
    h = pid_bh - b * H
    v0 = pid_v * BV
    qk_base = (b * T * H + h) * D
    s_base = (b * T * H + h) * BT

    state0 = tl.zeros((64, BV), tl.float32)
    state1 = tl.zeros((64, BV), tl.float32)

    chunk = 0
    while chunk < T // BT:
        tc = chunk * BT

        p_w0 = tl.make_block_ptr(w_ptr + qk_base, (T, D), (H * D, 1), (tc, 0), (BT, 64), (1, 0))
        p_w1 = tl.make_block_ptr(w_ptr + qk_base, (T, D), (H * D, 1), (tc, 64), (BT, 64), (1, 0))
        w0 = tl.load(p_w0)
        w1 = tl.load(p_w1)
        correction = tl.dot(w0, state0.to(tl.bfloat16))
        correction += tl.dot(w1, state1.to(tl.bfloat16))

        p_u = tl.make_block_ptr(u_ptr + qk_base, (T, D), (H * D, 1), (tc, v0), (BT, BV), (1, 0))
        value = tl.load(p_u).to(tl.float32) - correction
        value_bf = value.to(tl.bfloat16)

        p_q0 = tl.make_block_ptr(q_ptr + qk_base, (T, D), (H * D, 1), (tc, 0), (BT, 64), (1, 0))
        p_q1 = tl.make_block_ptr(q_ptr + qk_base, (T, D), (H * D, 1), (tc, 64), (BT, 64), (1, 0))
        p_g0 = tl.make_block_ptr(gc_ptr + qk_base, (T, D), (H * D, 1), (tc, 0), (BT, 64), (1, 0))
        p_g1 = tl.make_block_ptr(gc_ptr + qk_base, (T, D), (H * D, 1), (tc, 64), (BT, 64), (1, 0))
        g0 = tl.load(p_g0)
        g1 = tl.load(p_g1)
        q0 = (tl.load(p_q0).to(tl.float32) * tl.exp(g0)).to(tl.bfloat16)
        q1 = (tl.load(p_q1).to(tl.float32) * tl.exp(g1)).to(tl.bfloat16)
        output = tl.dot(q0, state0.to(tl.bfloat16))
        output += tl.dot(q1, state1.to(tl.bfloat16))
        output *= SCALE

        p_score = tl.make_block_ptr(
            score_ptr + s_base, (T, BT), (H * BT, 1), (tc, 0), (BT, BT), (1, 0)
        )
        score = tl.load(p_score)
        output += tl.dot(score, value_bf)
        p_out = tl.make_block_ptr(out_ptr + qk_base, (T, D), (H * D, 1), (tc, v0), (BT, BV), (1, 0))
        tl.store(p_out, output.to(tl.bfloat16))

        ik = tl.arange(0, 64)
        last0 = tl.load(gc_ptr + qk_base + (tc + BT - 1) * H * D + ik)
        last1 = tl.load(gc_ptr + qk_base + (tc + BT - 1) * H * D + 64 + ik)
        state0 *= tl.exp(last0)[:, None]
        state1 *= tl.exp(last1)[:, None]

        p_kg0 = tl.make_block_ptr(kg_ptr + qk_base, (D, T), (1, H * D), (0, tc), (64, BT), (0, 1))
        p_kg1 = tl.make_block_ptr(kg_ptr + qk_base, (D, T), (1, H * D), (64, tc), (64, BT), (0, 1))
        state0 += tl.dot(tl.load(p_kg0), value_bf)
        state1 += tl.dot(tl.load(p_kg1), value_bf)
        chunk += 1


@triton.jit
def _chunks_recurrence128_kernel(
    q_ptr,
    gc_ptr,
    score_ptr,
    w_ptr,
    u_ptr,
    kg_ptr,
    out_ptr,
    T,
    H: tl.constexpr,
    D: tl.constexpr,
    BT: tl.constexpr,
    BV: tl.constexpr,
    SCALE: tl.constexpr,
):
    pid_v = tl.program_id(0)
    pid_bh = tl.program_id(1)
    b = pid_bh // H
    h = pid_bh - b * H
    v0 = pid_v * BV
    qk_base = (b * T * H + h) * D
    s_base = (b * T * H + h) * BT
    state = tl.zeros((D, BV), tl.float32)

    chunk = 0
    while chunk < T // BT:
        tc = chunk * BT
        p_w = tl.make_block_ptr(w_ptr + qk_base, (T, D), (H * D, 1), (tc, 0), (BT, D), (1, 0))
        p_u = tl.make_block_ptr(u_ptr + qk_base, (T, D), (H * D, 1), (tc, v0), (BT, BV), (1, 0))
        value = tl.load(p_u).to(tl.float32) - tl.dot(tl.load(p_w), state.to(tl.bfloat16))
        value_bf = value.to(tl.bfloat16)

        p_q = tl.make_block_ptr(q_ptr + qk_base, (T, D), (H * D, 1), (tc, 0), (BT, D), (1, 0))
        p_g = tl.make_block_ptr(gc_ptr + qk_base, (T, D), (H * D, 1), (tc, 0), (BT, D), (1, 0))
        g = tl.load(p_g)
        qg = (tl.load(p_q).to(tl.float32) * tl.exp(g)).to(tl.bfloat16)
        output = tl.dot(qg, state.to(tl.bfloat16)) * SCALE
        p_score = tl.make_block_ptr(score_ptr + s_base, (T, BT), (H * BT, 1), (tc, 0), (BT, BT), (1, 0))
        output += tl.dot(tl.load(p_score), value_bf)
        p_out = tl.make_block_ptr(out_ptr + qk_base, (T, D), (H * D, 1), (tc, v0), (BT, BV), (1, 0))
        tl.store(p_out, output.to(tl.bfloat16))

        ik = tl.arange(0, D)
        last = tl.load(gc_ptr + qk_base + (tc + BT - 1) * H * D + ik)
        state *= tl.exp(last)[:, None]
        p_kg = tl.make_block_ptr(kg_ptr + qk_base, (D, T), (1, H * D), (0, tc), (D, BT), (0, 1))
        state += tl.dot(tl.load(p_kg), value_bf)
        chunk += 1


@triton.jit
def _state_chunks_kernel(
    gc_ptr,
    w_ptr,
    u_ptr,
    kg_ptr,
    states_ptr,
    values_ptr,
    T,
    H: tl.constexpr,
    D: tl.constexpr,
    BT: tl.constexpr,
    BV: tl.constexpr,
):
    pid_v = tl.program_id(0)
    pid_bh = tl.program_id(1)
    b = pid_bh // H
    h = pid_bh - b * H
    v0 = pid_v * BV
    nt = T // BT
    qk_base = (b * T * H + h) * D
    state_base = (b * nt * H + h) * D * D
    state = tl.zeros((D, BV), tl.float32)

    chunk = 0
    while chunk < nt:
        tc = chunk * BT
        p_state = tl.make_block_ptr(
            states_ptr + state_base + chunk * H * D * D,
            (D, D), (D, 1), (0, v0), (D, BV), (1, 0),
        )
        tl.store(p_state, state.to(tl.bfloat16))

        p_w = tl.make_block_ptr(w_ptr + qk_base, (T, D), (H * D, 1), (tc, 0), (BT, D), (1, 0))
        p_u = tl.make_block_ptr(u_ptr + qk_base, (T, D), (H * D, 1), (tc, v0), (BT, BV), (1, 0))
        value = tl.load(p_u).to(tl.float32) - tl.dot(tl.load(p_w), state.to(tl.bfloat16))
        value_bf = value.to(tl.bfloat16)
        p_value = tl.make_block_ptr(values_ptr + qk_base, (T, D), (H * D, 1), (tc, v0), (BT, BV), (1, 0))
        tl.store(p_value, value_bf)

        ik = tl.arange(0, D)
        last = tl.load(gc_ptr + qk_base + (tc + BT - 1) * H * D + ik)
        state *= tl.exp(last)[:, None]
        p_kg = tl.make_block_ptr(kg_ptr + qk_base, (D, T), (1, H * D), (0, tc), (D, BT), (0, 1))
        state += tl.dot(tl.load(p_kg), value_bf)
        chunk += 1


@triton.jit
def _output_chunks_kernel(
    q_ptr,
    gc_ptr,
    score_ptr,
    states_ptr,
    values_ptr,
    out_ptr,
    T,
    H: tl.constexpr,
    D: tl.constexpr,
    BT: tl.constexpr,
    BV: tl.constexpr,
    SCALE: tl.constexpr,
):
    pid_v = tl.program_id(0)
    chunk = tl.program_id(1)
    pid_bh = tl.program_id(2)
    b = pid_bh // H
    h = pid_bh - b * H
    v0 = pid_v * BV
    nt = T // BT
    tc = chunk * BT
    qk_base = (b * T * H + h) * D
    score_base = (b * T * H + h) * BT
    state_base = ((b * nt + chunk) * H + h) * D * D

    p_q = tl.make_block_ptr(q_ptr + qk_base, (T, D), (H * D, 1), (tc, 0), (BT, D), (1, 0))
    p_g = tl.make_block_ptr(gc_ptr + qk_base, (T, D), (H * D, 1), (tc, 0), (BT, D), (1, 0))
    p_state = tl.make_block_ptr(states_ptr + state_base, (D, D), (D, 1), (0, v0), (D, BV), (1, 0))
    qg = (tl.load(p_q).to(tl.float32) * tl.exp(tl.load(p_g))).to(tl.bfloat16)
    output = tl.dot(qg, tl.load(p_state)) * SCALE

    p_score = tl.make_block_ptr(score_ptr + score_base, (T, BT), (H * BT, 1), (tc, 0), (BT, BT), (1, 0))
    p_value = tl.make_block_ptr(values_ptr + qk_base, (T, D), (H * D, 1), (tc, v0), (BT, BV), (1, 0))
    output += tl.dot(tl.load(p_score), tl.load(p_value))
    p_out = tl.make_block_ptr(out_ptr + qk_base, (T, D), (H * D, 1), (tc, v0), (BT, BV), (1, 0))
    tl.store(p_out, output.to(tl.bfloat16))


def _forward(q, k, v, g, beta, scale: float):
    B, T, H, D = q.shape
    bt = 64
    nt = T // bt
    gc = torch.empty_like(g)
    score = torch.empty((B, T, H, bt), device=q.device, dtype=q.dtype)
    w = torch.empty_like(k)
    u = torch.empty_like(v)
    kg = torch.empty_like(k)
    out = torch.empty_like(v)

    lower = torch.empty((B, T, H, bt), device=q.device, dtype=torch.float32)
    transform = torch.empty_like(score)
    _pair_matrices_kernel[(nt, B * H)](
        q, k, g, gc, beta, score, lower,
        T, H=H, D=D, BT=bt, SCALE=scale,
        num_warps=PAIR_WARPS,
        num_stages=PAIR_STAGES,
    )
    _solve64_kernel[(nt, B * H)](
        lower, beta, transform, T, H=H, BT=bt, BC=16,
        num_warps=SOLVE_WARPS,
        num_stages=SOLVE_STAGES,
    )
    _project_chunks_kernel[(nt, B * H)](
        k, v, gc, transform, w, u, kg,
        T, H=H, D=D, BT=bt,
        num_warps=PROJECT_WARPS,
        num_stages=PROJECT_STAGES,
    )
    states = torch.empty((B, nt, H, D, D), device=q.device, dtype=q.dtype)
    values = torch.empty_like(v)
    long_context = T >= 4096
    state_bv = 32 if long_context else STATE_BV
    state_warps = 8 if long_context else STATE_WARPS
    state_stages = 4 if long_context else 5
    _state_chunks_kernel[(D // state_bv, B * H)](
        gc, w, u, kg, states, values,
        T, H=H, D=D, BT=bt, BV=state_bv,
        num_warps=state_warps,
        num_stages=state_stages,
    )
    output_bv = OUTPUT_BV
    _output_chunks_kernel[(D // output_bv, nt, B * H)](
        q, gc, score, states, values, out,
        T, H=H, D=D, BT=bt, BV=output_bv, SCALE=scale,
        num_warps=OUTPUT_WARPS,
        num_stages=OUTPUT_STAGES,
    )
    return out


class Model(nn.Module):
    def __init__(self, B: int, T: int, H: int, K: int, V: int, chunk_size: int = 64):
        super().__init__()
        if K != 128 or V != 128 or chunk_size != 64:
            raise ValueError("This kernel is specialized for K=V=128 and chunks of 64")
        self.scale = float(K) ** -0.5
        self.register_buffer("_dummy", torch.zeros(1), persistent=False)

    def forward(self, q, k, v, g, beta):
        return _forward(q, k, v, g, beta, self.scale)


B = 2
T = 1024
H = 8
K = 128
V = 128
CHUNK_SIZE = 64


def get_inputs():
    torch.manual_seed(0)
    q = torch.randn(B, T, H, K, dtype=torch.bfloat16) * 0.1
    k = torch.randn(B, T, H, K, dtype=torch.bfloat16) * 0.1
    v = torch.randn(B, T, H, V, dtype=torch.bfloat16) * 0.1
    g = torch.randn(B, T, H, K, dtype=torch.float32) * 0.1 - 0.05
    beta = torch.sigmoid(torch.randn(B, T, H, dtype=torch.bfloat16))
    return [q, k, v, g, beta]


def get_init_inputs():
    return [B, T, H, K, V, CHUNK_SIZE]

20260721_142442_codex_gpt-5.6-sol_02_kda_cutlass