KernelBench hard · RTX PRO 6000

KDA CUTLASS GLM-5.3

4.56%geomean peak fraction across shapes

manually audited: clean

Isolated-regrade 0.0456 on RTX PRO 6000. Custom KDA chunk forward: CUDA K1 (SM120 mma.sync via load_inline) plus a Triton two-level inter-chunk scan (k2a/k2b/k2c). Workspace is shape-keyed (B,T,H) scratch, rewritten every call; output is a fresh empty tensor. No FLA import. Lint CLEAN. template_mutated=false. Numeric stress on; check.py unmodified.

harnesszai-claudeagent session4h 48mtotal wall4h 59mcheck37sbenchmark2soutput tokens583,747cost$39.21regimecompute

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

2×1024×8×128×128×640.090 ms4.8%0.28 TB/s · 16% of 1.8 TB/s HBM · also 24 TFLOPS (5% of compute)
2×2048×8×128×128×640.146 ms5.9%0.35 TB/s · 19% of 1.8 TB/s HBM · also 29 TFLOPS (6% of compute)
1×4096×8×128×128×640.155 ms5.5%0.33 TB/s · 18% of 1.8 TB/s HBM · also 28 TFLOPS (6% of compute)
1×2048×4×128×128×640.078 ms2.8%0.16 TB/s · 9% of 1.8 TB/s HBM · also 14 TFLOPS (3% of compute)

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

geomean(4.8% · 5.9% · 5.5% · 2.8%) = 4.6%

Kernel source (redacted)
"""Kimi Delta Attention (KDA) forward, chunk form — custom Triton kernels for SM120.

Two-kernel decomposition. Per chunk (BT=64, K=V=128; T % 64 == 0 so no masks):

  K1 (grid NT x B*H, chunk-parallel)
    g  -> in-chunk fp32 cumsum
    A0 = (k e^{g-g_l}) (k e^{g_l-g})^T          (bf16 tensor-core dot)
    B  = tril(beta_row * A0, -1);  M = (I+B)^-1 via exact nilpotent doubling
    u  = M (beta v),   w = M (beta e^g k)
    P' = diag(e^{g_l}) - (k e^{g_l-g})^T w      [the chunk's state operator]
    C  = (k e^{g_l-g})^T u
    Aqk = scale * tril(q e^{g-g_l} (k e^{g_l-g})^T, incl diag)
    qe  = scale q e^g - Aqk w,   t = Aqk u
    stores: P', C (B,H,NT,128,128 bf16), qe, t (token-major bf16)

  K2 (grid V/BV x B*H, sequential over NT, fused output)
    S = C_i + P'_i @ S          (one dot per step; decay folded into P')
    o_i = qe_i @ S_i + t_i      (uses the pre-update state)

o = qg @ S_i + Aqk (u_i - w_i S_i) = (qg - Aqk w) S_i + Aqk u  (exact identity).
"""
from __future__ import annotations

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

RCP_LN2 = tl.constexpr(1.4426950408889634)


@triton.jit
def _col2(x, AX: tl.constexpr):
    """[AX, 128] -> two [AX, 64] column halves."""
    return tl.split(tl.permute(tl.reshape(x, (AX, 2, 64)), (0, 2, 1)))


@triton.jit
def _row2(x, AY: tl.constexpr):
    """[64, AY] -> two [32, AY] row halves."""
    return tl.split(tl.permute(tl.reshape(x, (2, 32, AY)), (1, 2, 0)))


@triton.jit
def _col2b(x, AX: tl.constexpr):
    """[AX, 64] -> two [AX, 32] column halves."""
    return tl.split(tl.permute(tl.reshape(x, (AX, 2, 32)), (0, 2, 1)))


# ---------------------------------------------------------------------------
# K1: chunk-local preparation (full-tile form, no layout shuffles)
# ---------------------------------------------------------------------------
@triton.jit
def _kda_k1(
    q, k, v, g, beta,
    P, C, Pd, qe, t,
    T,
    H: tl.constexpr,
    NT,
    scale,
):
    pid_t = tl.program_id(0)
    pid_bh = tl.program_id(1).to(tl.int64)
    i_b = pid_bh // H
    i_h = pid_bh - i_b * H

    o_c = tl.arange(0, 64)
    o_k = tl.arange(0, 128)
    tok0 = pid_t * 64
    rowb = (i_b * T * H + i_h)

    b_k = tl.load(k + rowb * 128 + (tok0 + o_c)[:, None] * (H * 128) + o_k[None, :])
    b_g = tl.load(g + rowb * 128 + (tok0 + o_c)[:, None] * (H * 128) + o_k[None, :])
    b_beta = tl.load(beta + rowb + (tok0 + o_c) * H).to(tl.float32)

    # cumsum via triangular-ones dot (tf32) instead of 6 shuffle rounds
    ltri = tl.where(o_c[:, None] >= o_c[None, :], 1.0, 0.0)
    b_g = tl.dot(ltri, b_g * RCP_LN2, input_precision="tf32")
    b_glast = tl.sum(tl.where((o_c[:, None] == 63), b_g, 0.0), 0)
    b_d = tl.exp2(b_glast)

    e_up = tl.exp2(b_g - b_glast[None, :])
    e_all = e_up * b_d[None, :]
    b_kf = b_k.to(tl.float32)
    b_kl = (b_kf * e_up).to(tl.bfloat16)
    b_kr = (b_kf * tl.exp2(b_glast[None, :] - b_g)).to(tl.bfloat16)

    # ---- A0, B; M = (I+B)^-1 by bf16 doubling ----
    b_B = tl.dot(b_kl, tl.trans(b_kr))
    b_B = tl.where(o_c[:, None] > o_c[None, :], b_B * b_beta[:, None], 0.0).to(tl.bfloat16)

    eye = tl.where(o_c[:, None] == o_c[None, :], 1.0, 0.0).to(tl.bfloat16)
    X = (eye.to(tl.float32) - b_B.to(tl.float32)).to(tl.bfloat16)
    b_B = tl.dot(b_B, b_B).to(tl.bfloat16)        # B^2
    X = tl.dot(X, (eye.to(tl.float32) + b_B.to(tl.float32)).to(tl.bfloat16)).to(tl.bfloat16)
    b_B = tl.dot(b_B, b_B).to(tl.bfloat16)        # B^4
    X = tl.dot(X, (eye.to(tl.float32) + b_B.to(tl.float32)).to(tl.bfloat16)).to(tl.bfloat16)
    b_B = tl.dot(b_B, b_B).to(tl.bfloat16)        # B^8
    X = tl.dot(X, (eye.to(tl.float32) + b_B.to(tl.float32)).to(tl.bfloat16)).to(tl.bfloat16)
    b_B = tl.dot(b_B, b_B).to(tl.bfloat16)        # B^16
    X = tl.dot(X, (eye.to(tl.float32) + b_B.to(tl.float32)).to(tl.bfloat16)).to(tl.bfloat16)
    b_B = tl.dot(b_B, b_B).to(tl.bfloat16)        # B^32
    M = tl.dot(X, (eye.to(tl.float32) + b_B.to(tl.float32)).to(tl.bfloat16)).to(tl.bfloat16)

    # ---- w = M kb, u = M vb ----
    b_v = tl.load(v + rowb * 128 + (tok0 + o_c)[:, None] * (H * 128) + o_k[None, :])
    b_kb = (b_kf * (b_beta[:, None] * e_all)).to(tl.bfloat16)
    b_vb = (b_v.to(tl.float32) * b_beta[:, None]).to(tl.bfloat16)
    b_w = tl.dot(M, b_kb).to(tl.bfloat16)
    b_u = tl.dot(M, b_vb).to(tl.bfloat16)

    # ---- P'' = -kr^T w, C = kr^T u ----
    pbase = (pid_bh * NT + pid_t) * (128 * 128)
    b_krt = tl.trans(b_kr)
    tl.store(P + pbase + o_k[:, None] * 128 + o_k[None, :], (-tl.dot(b_krt, b_w)).to(tl.bfloat16))
    tl.store(C + pbase + o_k[:, None] * 128 + o_k[None, :], tl.dot(b_krt, b_u).to(tl.bfloat16))
    tl.store(Pd + (pid_bh * NT + pid_t) * 128 + o_k, b_d)

    # ---- Aqk, qe, t ----
    b_q = tl.load(q + rowb * 128 + (tok0 + o_c)[:, None] * (H * 128) + o_k[None, :])
    b_ql = (b_q.to(tl.float32) * (scale * e_up)).to(tl.bfloat16)
    b_qg = (b_q.to(tl.float32) * (scale * e_all)).to(tl.bfloat16)
    b_Aqk = tl.dot(b_ql, tl.trans(b_kr))
    b_Aqk = tl.where(o_c[:, None] >= o_c[None, :], b_Aqk, 0.0)
    b_Aqk = b_Aqk.to(tl.bfloat16)
    b_qe = (b_qg.to(tl.float32) - tl.dot(b_Aqk, b_w)).to(tl.bfloat16)
    b_t = tl.dot(b_Aqk, b_u).to(tl.bfloat16)

    tl.store(qe + rowb * 128 + (tok0 + o_c)[:, None] * (H * 128) + o_k[None, :], b_qe)
    tl.store(t + rowb * 128 + (tok0 + o_c)[:, None] * (H * 128) + o_k[None, :], b_t)


# ---------------------------------------------------------------------------
# K2 pass A: per-group local scan (S=0 init) + group composite A_g
#   grid (NV, BH, G), BV columns each; A maintained full-width redundantly
# ---------------------------------------------------------------------------
@triton.jit
def _kda_k2a(
    P, C, Pd, A, Bg,
    T,
    H: tl.constexpr,
    NT,
    L: tl.constexpr,
    BV: tl.constexpr,
):
    i_v = tl.program_id(0)
    pid_bh = tl.program_id(1).to(tl.int64)
    i_g = tl.program_id(2)
    i_b = pid_bh // H
    i_h = pid_bh - i_b * H
    o_k = tl.arange(0, 128)
    o_v = i_v * BV + tl.arange(0, BV)
    rowb = (i_b * T * H + i_h)

    b_S = tl.zeros([128, BV], dtype=tl.float32)
    b_A = tl.where(o_k[:, None] == o_k[None, :], 1.0, 0.0)               # [128,128] fp32

    for j in range(L):
        i_t = i_g * L + j
        pbase = (pid_bh * NT + i_t) * (128 * 128)
        b_P = tl.load(P + pbase + o_k[:, None] * 128 + o_k[None, :])
        b_C = tl.load(C + pbase + o_k[:, None] * 128 + o_v[None, :])
        b_d = tl.load(Pd + (pid_bh * NT + i_t) * 128 + o_k)
        b_S = b_d[:, None] * b_S + b_C.to(tl.float32) + tl.dot(b_P, b_S.to(tl.bfloat16))
        b_A = b_d[:, None] * b_A + tl.dot(b_P, b_A.to(tl.bfloat16))

    abase = ((pid_bh * tl.num_programs(2)) + i_g) * (128 * 128)
    tl.store(A + abase + o_k[:, None] * 128 + o_k[None, :], b_A.to(tl.bfloat16))
    tl.store(Bg + abase + o_k[:, None] * 128 + o_v[None, :], b_S.to(tl.bfloat16))


# ---------------------------------------------------------------------------
# K2 pass B: tiny sequential scan over the G group composites -> boundary states
#   grid (NV, BH)
# ---------------------------------------------------------------------------
@triton.jit
def _kda_k2b(
    A, Bg, Sb,
    H: tl.constexpr,
    G: tl.constexpr,
    BV: tl.constexpr,
):
    i_v = tl.program_id(0)
    pid_bh = tl.program_id(1).to(tl.int64)
    o_k = tl.arange(0, 128)
    o_v = i_v * BV + tl.arange(0, BV)

    b_S = tl.zeros([128, BV], dtype=tl.float32)
    for i_g in range(G):
        abase = ((pid_bh * G) + i_g) * (128 * 128)
        b_A = tl.load(A + abase + o_k[:, None] * 128 + o_k[None, :])
        b_Bg = tl.load(Bg + abase + o_k[:, None] * 128 + o_v[None, :])
        tl.store(Sb + abase + o_k[:, None] * 128 + o_v[None, :], b_S.to(tl.bfloat16))
        b_S = b_Bg.to(tl.float32) + tl.dot(b_A, b_S.to(tl.bfloat16))


# ---------------------------------------------------------------------------
# K2 pass C: final scan within each group from its boundary state + outputs
#   grid (NV, BH, G)
# ---------------------------------------------------------------------------
@triton.jit
def _kda_k2c(
    P, C, Pd, qe, t, Sb,
    o,
    T,
    H: tl.constexpr,
    NT,
    L: tl.constexpr,
    BV: tl.constexpr,
):
    i_v = tl.program_id(0)
    pid_bh = tl.program_id(1).to(tl.int64)
    i_g = tl.program_id(2)
    i_b = pid_bh // H
    i_h = pid_bh - i_b * H
    o_c = tl.arange(0, 64)
    o_k = tl.arange(0, 128)
    o_v = i_v * BV + tl.arange(0, BV)
    rowb = (i_b * T * H + i_h)

    abase = ((pid_bh * tl.num_programs(2)) + i_g) * (128 * 128)
    b_S = tl.load(Sb + abase + o_k[:, None] * 128 + o_v[None, :]).to(tl.float32)

    for j in range(L):
        i_t = i_g * L + j
        tok0 = i_t * 64
        pbase = (pid_bh * NT + i_t) * (128 * 128)
        b_P = tl.load(P + pbase + o_k[:, None] * 128 + o_k[None, :])
        b_C = tl.load(C + pbase + o_k[:, None] * 128 + o_v[None, :])
        b_d = tl.load(Pd + (pid_bh * NT + i_t) * 128 + o_k)
        b_qe = tl.load(qe + rowb * 128 + (tok0 + o_c)[:, None] * (H * 128) + o_k[None, :])
        b_t = tl.load(t + rowb * 128 + (tok0 + o_c)[:, None] * (H * 128) + o_v[None, :])

        b_o = tl.dot(b_qe, b_S.to(tl.bfloat16)) + b_t.to(tl.float32)
        tl.store(o + rowb * 128 + (tok0 + o_c)[:, None] * (H * 128) + o_v[None, :],
                 b_o.to(tl.bfloat16))
        b_S = b_d[:, None] * b_S + b_C.to(tl.float32) + tl.dot(b_P, b_S.to(tl.bfloat16))


# ---------------------------------------------------------------------------
# K2 (fallback): single sequential scan with fused output
# ---------------------------------------------------------------------------
@triton.jit
def _kda_k2(
    P, C, Pd, qe, t,
    o,
    T,
    H: tl.constexpr,
    NT,
    BV: tl.constexpr,
):
    i_v = tl.program_id(0)
    pid_bh = tl.program_id(1).to(tl.int64)
    i_b = pid_bh // H
    i_h = pid_bh - i_b * H

    o_c = tl.arange(0, 64)
    o_k = tl.arange(0, 128)
    o_v = i_v * BV + tl.arange(0, BV)
    rowb = (i_b * T * H + i_h)

    b_S = tl.zeros([128, BV], dtype=tl.float32)

    for i_t in range(NT):
        tok0 = i_t * 64
        pbase = (pid_bh * NT + i_t) * (128 * 128)
        b_P = tl.load(P + pbase + o_k[:, None] * 128 + o_k[None, :])       # [128,128]
        b_C = tl.load(C + pbase + o_k[:, None] * 128 + o_v[None, :])       # [128,BV]
        b_d = tl.load(Pd + (pid_bh * NT + i_t) * 128 + o_k)                # [128] fp32
        b_qe = tl.load(qe + rowb * 128 + (tok0 + o_c)[:, None] * (H * 128) + o_k[None, :])
        b_t = tl.load(t + rowb * 128 + (tok0 + o_c)[:, None] * (H * 128) + o_v[None, :])

        b_o = tl.dot(b_qe, b_S.to(tl.bfloat16)) + b_t.to(tl.float32)       # pre-update S
        tl.store(o + rowb * 128 + (tok0 + o_c)[:, None] * (H * 128) + o_v[None, :],
                 b_o.to(tl.bfloat16))
        b_S = b_d[:, None] * b_S + b_C.to(tl.float32) + tl.dot(b_P, b_S.to(tl.bfloat16))


# ---------------------------------------------------------------------------
# Model
# ---------------------------------------------------------------------------
def _kda_forward(q, k, v, g, beta, scale, chunk_size=64):
    B, T, H, K = q.shape
    V = v.shape[-1]
    assert K == 128 and V == 128 and chunk_size == 64 and T % 64 == 0
    NT = T // 64
    BH = B * H
    dev = q.device

    bufs = _BUF_CACHE.get((B, T, H))
    if bufs is None:
        # two-level scan split: aim for ~128 CTAs in k2a/k2c (2*BV-halves per bh)
        tgt = max(1, 128 // (2 * BH))
        G = 1
        while G * 2 <= tgt and NT % (G * 2) == 0:
            G *= 2
        bufs = (
            torch.empty(B, H, NT, K, V, dtype=torch.bfloat16, device=dev),  # P
            torch.empty(B, H, NT, K, V, dtype=torch.bfloat16, device=dev),  # C
            torch.empty(B * H * NT, K, dtype=torch.float32, device=dev),    # Pd
            torch.empty(B, T, H, K, dtype=torch.bfloat16, device=dev),      # qe
            torch.empty(B, T, H, V, dtype=torch.bfloat16, device=dev),      # t
            torch.empty(B * H * G, K, V, dtype=torch.bfloat16, device=dev),  # A
            torch.empty(B * H * G, K, V, dtype=torch.bfloat16, device=dev),  # Bg
            torch.empty(B * H * G, K, V, dtype=torch.bfloat16, device=dev),  # Sb
            G,
        )
        _BUF_CACHE[(B, T, H)] = bufs
    P_, C_, Pd_, qe_, t_, A_, Bg_, Sb_, G = bufs
    L = NT // G
    o = torch.empty(B, T, H, V, dtype=torch.bfloat16, device=dev)

    _run_k1(q, k, v, g, beta, P_, C_, Pd_, qe_, t_, T, H, NT, scale, BH)
    _kda_k2a[(2, BH, G)](P_, C_, Pd_, A_, Bg_, T, H, NT, L=L, BV=64,
                         num_warps=8, num_stages=2)
    _kda_k2b[(2, BH)](A_, Bg_, Sb_, H, G=G, BV=64, num_warps=8, num_stages=2)
    _kda_k2c[(2, BH, G)](P_, C_, Pd_, qe_, t_, Sb_, o, T, H, NT, L=L, BV=64,
                         num_warps=8, num_stages=2)
    return o




# ---------------- CUDA K1 (SM120 mma.sync path) ----------------
_CUDA_K1_SRC = r"""
#include <torch/extension.h>
#include <cuda_bf16.h>

#define DEV __device__ __forceinline__

DEV int lane_id() { return threadIdx.x & 31; }
DEV uint32_t smem_u32(const void* p) { return (uint32_t)__cvta_generic_to_shared(p); }

DEV void ldm_x4(uint32_t& r0, uint32_t& r1, uint32_t& r2, uint32_t& r3, const void* p) {
    asm volatile("ldmatrix.sync.aligned.m8n8.x4.shared.b16 {%0,%1,%2,%3}, [%4];\n"
                 : "=r"(r0), "=r"(r1), "=r"(r2), "=r"(r3) : "r"(smem_u32(p)));
}
DEV void ldm_x4t(uint32_t& r0, uint32_t& r1, uint32_t& r2, uint32_t& r3, const void* p) {
    asm volatile("ldmatrix.sync.aligned.m8n8.x4.trans.shared.b16 {%0,%1,%2,%3}, [%4];\n"
                 : "=r"(r0), "=r"(r1), "=r"(r2), "=r"(r3) : "r"(smem_u32(p)));
}
DEV void stm_x4(const void* p, uint32_t r0, uint32_t r1, uint32_t r2, uint32_t r3) {
    asm volatile("stmatrix.sync.aligned.m8n8.x4.shared.b16 [%0], {%1,%2,%3,%4};\n"
                 :: "r"(smem_u32(p)), "r"(r0), "r"(r1), "r"(r2), "r"(r3));
}
DEV void mma_bf16(float& d0, float& d1, float& d2, float& d3,
                  uint32_t a0, uint32_t a1, uint32_t a2, uint32_t a3, uint32_t b0, uint32_t b1) {
    asm volatile("mma.sync.aligned.m16n8k16.row.col.f32.bf16.bf16.f32 "
                 "{%0,%1,%2,%3}, {%4,%5,%6,%7}, {%8,%9}, {%0,%1,%2,%3};\n"
                 : "+f"(d0), "+f"(d1), "+f"(d2), "+f"(d3)
                 : "r"(a0), "r"(a1), "r"(a2), "r"(a3), "r"(b0), "r"(b1));
}
DEV uint32_t pack_bf16(float lo, float hi) {
    __nv_bfloat162 v = __floats2bfloat162_rn(lo, hi);
    return *reinterpret_cast<uint32_t*>(&v);
}
DEV float bf2f(__nv_bfloat16 x) { return __bfloat162float(x); }

#ifdef K1_DBG
__device__ bool k1_dbg_root() { return blockIdx.x == 0 && blockIdx.y == 0; }
DEV void dump_f32(float* dbg, long base, const float* buf, int n) {
    if (blockIdx.x != 0 || blockIdx.y != 0) return;
    for (int i = threadIdx.x; i < n; i += 256) dbg[base + i] = buf[i];
}
DEV void dump_bf(float* dbg, long base, const __nv_bfloat16* buf, int n) {
    if (!k1_dbg_root()) return;
    for (int i = threadIdx.x; i < n; i += 256) dbg[base + i] = bf2f(buf[i]);
}
#endif

// row pitch must be a multiple of 8 elems (16B units): S % 8 == 0 makes the XOR
// swizzle conflict-free for ldmatrix/stmatrix 8-row groups (unit%S differs only via q^r).
constexpr int LD128 = 128;
constexpr int LD64  = 64;
constexpr float RCP_LN2 = 1.4426950408889634f;

struct Smem {
    __nv_bfloat16 kl[64 * LD128];   // k*e_up (A0 operand) -> qe/t/P/C store staging
    __nv_bfloat16 qlq[64 * LD128];  // q*scale*e_up (Aqk operand), written in prologue
    __nv_bfloat16 krt[128 * LD64];  // kr^T row-major [k=128][c=64]
    __nv_bfloat16 kb[64 * LD128];   // k*beta*e_all -> w
    __nv_bfloat16 vb[64 * LD128];   // v*beta -> u
    __nv_bfloat16 ba[64 * LD64];    // chain ping / M
    __nv_bfloat16 bb[64 * LD64];    // chain pong / Aqk
    float beta_s[64];
    float d_s[128];
    float part[512];
};

// swizzled offsets: XOR the 8-element chunk index with (row & 7)
DEV int o128(int r, int c) { return r * LD128 + ((c >> 3) ^ (r & 7)) * 8 + (c & 7); }
DEV int o64(int r, int c)  { return r * LD64 + ((c >> 3) ^ (r & 7)) * 8 + (c & 7); }

// acc[RM/16][RN/8][4] += A @ B for this warp's [RM, RN] slice.
//   A: row-major [M,K] smem (swizzled), warp rows arow0..+RM, row pitch aw
//   B: row-major [K,N] smem (swizzled), warp cols bcol0..+RN, row pitch bw
template <int RM, int RN>
DEV void dot_mn(float acc[RM / 16][RN / 8][4], const __nv_bfloat16* A, int arow0, int aw,
                const __nv_bfloat16* B, int bcol0, int bw, int K) {
    int l = lane_id();
    for (int k0 = 0; k0 < K; k0 += 16) {
        uint32_t a[RM / 16][4];
#pragma unroll
        for (int mi = 0; mi < RM / 16; mi++) {
            int r = arow0 + mi * 16 + (l % 16);
            int c = k0 + (l / 16) * 8;
            ldm_x4(a[mi][0], a[mi][1], a[mi][2], a[mi][3],
                   A + (aw == LD128 ? o128(r, c) : o64(r, c)));
        }
        uint32_t b[RN / 4];   // each x4t yields 2 frags x 2 regs
#pragma unroll
        for (int np = 0; np < RN / 16; np++) {
            uint32_t r0, r1, r2, r3;
            int g = l / 8;
            int kk = k0 + (l % 8) + (g & 1) * 8;
            int nn = bcol0 + np * 16 + (g >> 1) * 8;
            ldm_x4t(r0, r1, r2, r3, B + (bw == LD128 ? o128(kk, nn) : o64(kk, nn)));
            b[np * 4 + 0] = r0; b[np * 4 + 1] = r1;
            b[np * 4 + 2] = r2; b[np * 4 + 3] = r3;
        }
#pragma unroll
        for (int mi = 0; mi < RM / 16; mi++)
#pragma unroll
            for (int ni = 0; ni < RN / 8; ni++)
                mma_bf16(acc[mi][ni][0], acc[mi][ni][1], acc[mi][ni][2], acc[mi][ni][3],
                         a[mi][0], a[mi][1], a[mi][2], a[mi][3], b[ni * 2], b[ni * 2 + 1]);
    }
}

// zero acc
template <int RM, int RN>
DEV void zero_acc(float acc[RM / 16][RN / 8][4]) {
#pragma unroll
    for (int i = 0; i < RM / 16; i++)
#pragma unroll
        for (int j = 0; j < RN / 8; j++)
#pragma unroll
            for (int p = 0; p < 4; p++) acc[i][j][p] = 0.f;
}

// stage warp's [RM, RN] acc slice to a 64- or 128-wide smem tile at (row0, col0).
// mode: 0 plain, 1 negate.  LD picks o64/o128 addressing.
template <int RM, int RN, int LD>
DEV void stage_acc(float acc[RM / 16][RN / 8][4], int row0, int col0, __nv_bfloat16* base,
                   bool negate) {
    int l = lane_id();
#pragma unroll
    for (int mi = 0; mi < RM / 16; mi++)
#pragma unroll
        for (int ni = 0; ni < RN / 8; ni += 2) {
            int r = row0 + mi * 16;
            int c = col0 + ni * 8;
            int lr = r + (l % 8) + ((l / 8) & 1) * 8;
            int lc = c + (l / 16) * 8;
            int off = (LD == LD128) ? o128(lr, lc) : o64(lr, lc);
            float s0 = negate ? -acc[mi][ni][0] : acc[mi][ni][0];
            float s1 = negate ? -acc[mi][ni][1] : acc[mi][ni][1];
            float s2 = negate ? -acc[mi][ni][2] : acc[mi][ni][2];
            float s3 = negate ? -acc[mi][ni][3] : acc[mi][ni][3];
            float u0 = negate ? -acc[mi][ni + 1][0] : acc[mi][ni + 1][0];
            float u1 = negate ? -acc[mi][ni + 1][1] : acc[mi][ni + 1][1];
            float u2 = negate ? -acc[mi][ni + 1][2] : acc[mi][ni + 1][2];
            float u3 = negate ? -acc[mi][ni + 1][3] : acc[mi][ni + 1][3];
            stm_x4(base + off, pack_bf16(s0, s1), pack_bf16(s2, s3),
                   pack_bf16(u0, u1), pack_bf16(u2, u3));
        }
}

// copy a [64,128] bf16 smem tile (swizzled, pitch LD128) to gmem rows rowG..rowG+63,
// gmem row pitch = gstride elements, i.e. dst + (rowG+i)*gstride + j.
DEV void tile_to_gmem(const __nv_bfloat16* src, __nv_bfloat16* dst, int rowG, long gstride) {
    int t = threadIdx.x;
    int r = t >> 2;
    int u = (t & 3) * 4;                 // uint4 index (8 bf16 each)
    const uint4* s4 = reinterpret_cast<const uint4*>(src);
    uint4* d4 = reinterpret_cast<uint4*>(dst);
    long grow = (long)(rowG + r) * (gstride >> 3);   // gstride in bf16 -> in uint4
#pragma unroll
    for (int i = 0; i < 4; i++) {
        int uu = u + i;
        int soff = o128(r, uu * 8) >> 3;            // /8 elems -> /8 in uint4? (8 bf16 = 1 uint4) but swizzle offset is in elems: elem offset o128(r, uu*8) is a multiple of 8 -> uint4 index = /8
        d4[grow + uu] = s4[soff];
    }
}

// ---------------------------------------------------------------------------
// K1: grid (NT, BH), 256 threads (8 warps: wm = w>>2, wn = w&3).
// ---------------------------------------------------------------------------
__global__ void __launch_bounds__(256, 1) kda_k1(
    const __nv_bfloat16* __restrict__ q,
    const __nv_bfloat16* __restrict__ k,
    const __nv_bfloat16* __restrict__ v,
    const float* __restrict__ g,
    const __nv_bfloat16* __restrict__ beta,
    __nv_bfloat16* __restrict__ P,
    __nv_bfloat16* __restrict__ Cm,
    float* __restrict__ Pd,
    __nv_bfloat16* __restrict__ qe,
    __nv_bfloat16* __restrict__ tt,
    long T, int H, long NT, float scale
#ifdef K1_DBG
    , float* dbg
#endif
) {
    extern __shared__ __nv_bfloat16 smem_raw[];
    Smem& s = *reinterpret_cast<Smem*>(smem_raw);
    int pid_t = blockIdx.x;
    long bh = blockIdx.y;
    long i_b = bh / H, i_h = bh - i_b * H;
    long tok0 = (long)pid_t * 64;

    int tid = threadIdx.x, l = lane_id(), w = tid >> 5;
    int wm = w >> 2, wn = w & 3;

    long HK = (long)H * 128;
    long rowK = (i_b * T + tok0) * HK + i_h * 128;      // k/v/q/g row (t-major, K-contig)
    long rowB = (i_b * T + tok0) * H + i_h;             // beta row
    long outQ = (i_b * T + tok0) * HK + i_h * 128;      // qe/t rows

    // ---------------- prologue ----------------
    // thread owns 4 columns c0..c0+3 and 8 rows r0..r0+7 (rg = tid>>5, cp = tid&31).
    // All gmem loads vectorized: g float4, k/v/q 8B. part[] overlays s.ba (free here).
    {
        int cp = tid & 31;
        int c0 = cp * 4;
        int rg = tid >> 5;
        int r0 = rg * 8;
        const float4* gp = reinterpret_cast<const float4*>(g + rowK + (long)r0 * HK + c0);
        const uint2* kp = reinterpret_cast<const uint2*>(k + rowK + (long)r0 * HK + c0);
        const uint2* vp = reinterpret_cast<const uint2*>(v + rowK + (long)r0 * HK + c0);
        const uint2* qp = reinterpret_cast<const uint2*>(q + rowK + (long)r0 * HK + c0);
        float4 gx[8]; uint2 kr4[8], vr4[8], qr4[8];
#pragma unroll
        for (int i = 0; i < 8; i++) gx[i] = gp[(long)i * (HK >> 2)];
#pragma unroll
        for (int i = 0; i < 8; i++) kr4[i] = kp[(long)i * (HK >> 2)];
#pragma unroll
        for (int i = 0; i < 8; i++) vr4[i] = vp[(long)i * (HK >> 2)];
#pragma unroll
        for (int i = 0; i < 8; i++) qr4[i] = qp[(long)i * (HK >> 2)];
        float acc[8][4];
        float run[4] = {0.f, 0.f, 0.f, 0.f};
#pragma unroll
        for (int i = 0; i < 8; i++) {
            run[0] += gx[i].x * RCP_LN2; acc[i][0] = run[0];
            run[1] += gx[i].y * RCP_LN2; acc[i][1] = run[1];
            run[2] += gx[i].z * RCP_LN2; acc[i][2] = run[2];
            run[3] += gx[i].w * RCP_LN2; acc[i][3] = run[3];
        }
        float* part = reinterpret_cast<float*>(s.ba);   // [8][128], ba not used yet
#pragma unroll
        for (int j = 0; j < 4; j++) part[rg * 128 + c0 + j] = run[j];
        if (tid < 64) s.beta_s[tid] = bf2f(beta[rowB + (long)tid * H]);
        __syncthreads();
        float glast[4], pbelow[4];
#pragma unroll
        for (int j = 0; j < 4; j++) {
            float tot = 0.f, below = 0.f;
            for (int rr = 0; rr < 8; rr++) {
                float p = part[rr * 128 + c0 + j];
                tot += p;
                if (rr < rg) below += p;
            }
            glast[j] = tot;
            pbelow[j] = below;
        }
        if (rg == 0) {
            float4 dd = make_float4(exp2f(glast[0]), exp2f(glast[1]),
                                    exp2f(glast[2]), exp2f(glast[3]));
            reinterpret_cast<float4*>(s.d_s)[cp] = dd;
            reinterpret_cast<float4*>(Pd + (bh * NT + pid_t) * 128)[cp] = dd;
        }
        __nv_bfloat162* kl2 = reinterpret_cast<__nv_bfloat162*>(s.kl);
        __nv_bfloat162* ql2 = reinterpret_cast<__nv_bfloat162*>(s.qlq);
        __nv_bfloat162* kb2 = reinterpret_cast<__nv_bfloat162*>(s.kb);
        __nv_bfloat162* vb2 = reinterpret_cast<__nv_bfloat162*>(s.vb);
#pragma unroll
        for (int i = 0; i < 8; i++) {
            int r = r0 + i;
            __nv_bfloat162 k01 = *reinterpret_cast<__nv_bfloat162*>(&kr4[i]);
            __nv_bfloat162 k23 = *reinterpret_cast<__nv_bfloat162*>(&kr4[i].y);
            __nv_bfloat162 v01 = *reinterpret_cast<__nv_bfloat162*>(&vr4[i]);
            __nv_bfloat162 v23 = *reinterpret_cast<__nv_bfloat162*>(&vr4[i].y);
            __nv_bfloat162 q01 = *reinterpret_cast<__nv_bfloat162*>(&qr4[i]);
            __nv_bfloat162 q23 = *reinterpret_cast<__nv_bfloat162*>(&qr4[i].y);
            float bt = s.beta_s[r];
            float e_up[4], e_dn[4], ea[4];
#pragma unroll
            for (int j = 0; j < 4; j++) {
                float gc = acc[i][j] + pbelow[j];
                e_up[j] = exp2f(gc - glast[j]);
                e_dn[j] = exp2f(glast[j] - gc);
                ea[j]   = exp2f(gc);
            }
            int o = o128(r, c0) >> 1;
            kl2[o]     = __floats2bfloat162_rn(bf2f(k01.x) * e_up[0], bf2f(k01.y) * e_up[1]);
            kl2[o + 1] = __floats2bfloat162_rn(bf2f(k23.x) * e_up[2], bf2f(k23.y) * e_up[3]);
            ql2[o]     = __floats2bfloat162_rn(bf2f(q01.x) * scale * e_up[0], bf2f(q01.y) * scale * e_up[1]);
            ql2[o + 1] = __floats2bfloat162_rn(bf2f(q23.x) * scale * e_up[2], bf2f(q23.y) * scale * e_up[3]);
            kb2[o]     = __floats2bfloat162_rn(bf2f(k01.x) * bt * ea[0], bf2f(k01.y) * bt * ea[1]);
            kb2[o + 1] = __floats2bfloat162_rn(bf2f(k23.x) * bt * ea[2], bf2f(k23.y) * bt * ea[3]);
            vb2[o]     = __floats2bfloat162_rn(bf2f(v01.x) * bt, bf2f(v01.y) * bt);
            vb2[o + 1] = __floats2bfloat162_rn(bf2f(v23.x) * bt, bf2f(v23.y) * bt);
            s.krt[o64(c0, r)]     = __float2bfloat16(bf2f(k01.x) * e_dn[0]);
            s.krt[o64(c0 + 1, r)] = __float2bfloat16(bf2f(k01.y) * e_dn[1]);
            s.krt[o64(c0 + 2, r)] = __float2bfloat16(bf2f(k23.x) * e_dn[2]);
            s.krt[o64(c0 + 3, r)] = __float2bfloat16(bf2f(k23.y) * e_dn[3]);
        }
        __syncthreads();
    }
#ifdef K1_DBG
    dump_bf(dbg, 0, s.kl, 64*LD128);
    dump_bf(dbg, 61440, s.qlq, 64*LD128);
    dump_f32(dbg, 70000, s.beta_s, 64);
    dump_bf(dbg, 8192, s.krt, 128*LD64);
    dump_bf(dbg, 16384, s.kb, 64*LD128);
    dump_bf(dbg, 24576, s.vb, 64*LD128);
    __syncthreads();
#endif

    // ---------------- A0 = kl @ krT ; B = tril(A0*beta, -1) ; X = I - B ----------------
#ifdef K1_EARLY
    if (K1_EARLY == 1) return;
#endif
    float a0[2][2][4];
    zero_acc<32, 16>(a0);
    dot_mn<32, 16>(a0, s.kl, wm * 32, LD128, s.krt, wn * 16, LD64, 128);
    __syncthreads();                       // kl free after all warps' A0
#pragma unroll
    for (int mi = 0; mi < 2; mi++)
#pragma unroll
        for (int ni = 0; ni < 2; ni++)
#pragma unroll
            for (int p = 0; p < 4; p++) {
                int r, c;
                // frag coords within [32,16] warp slice
                r = wm * 32 + mi * 16 + (l >> 2);
                c = wn * 16 + ni * 8 + ((l & 3) * 2 + (p & 1));
                if (p >= 2) r += 8;
                float val = (r > c) ? a0[mi][ni][p] * s.beta_s[r] : 0.f;
                a0[mi][ni][p] = val;
            }
    float xacc[2][2][4];
#pragma unroll
    for (int mi = 0; mi < 2; mi++)
#pragma unroll
        for (int ni = 0; ni < 2; ni++)
#pragma unroll
            for (int p = 0; p < 4; p++) {
                int r = wm * 32 + mi * 16 + (l >> 2) + (p >= 2 ? 8 : 0);
                int c = wn * 16 + ni * 8 + ((l & 3) * 2 + (p & 1));
                xacc[mi][ni][p] = (r == c) ? 1.f : -a0[mi][ni][p];
            }
#ifdef K1_DBG
    if (blockIdx.x == 0 && blockIdx.y == 0 && w == 0)
        for (int mi = 0; mi < 2; mi++) for (int ni = 0; ni < 2; ni++) for (int p = 0; p < 4; p++)
            dbg[71000 + mi * 8 + ni * 4 + p] = a0[mi][ni][p];
    __syncthreads();
#endif
    stage_acc<32, 16, LD64>(a0, wm * 32, wn * 16, s.ba, false);
    __syncthreads();
#ifdef K1_DBG
    dump_bf(dbg, 32768, s.ba, 64*LD64);
    __syncthreads();
#endif

#ifdef K1_EARLY
    if (K1_EARLY == 2) return;
#endif
    // ---------------- chain: X <- X + X @ B^(2^i) ----------------
    __nv_bfloat16* cur = s.ba;
    __nv_bfloat16* oth = s.bb;
#pragma unroll
    for (int it = 0; it < 5; it++) {
        float sq[2][2][4];
        zero_acc<32, 16>(sq);
        dot_mn<32, 16>(sq, cur, wm * 32, LD64, cur, wn * 16, LD64, 64);
        stage_acc<32, 16, LD64>(sq, wm * 32, wn * 16, oth, false);
        __syncthreads();                  // cur dead: all warps done squaring
        stage_acc<32, 16, LD64>(xacc, wm * 32, wn * 16, cur, false);
        __syncthreads();
        dot_mn<32, 16>(xacc, cur, wm * 32, LD64, oth, wn * 16, LD64, 64);   // acc-init = X
        __syncthreads();
        __nv_bfloat16* tmp = cur; cur = oth; oth = tmp;
    }
    // M = xacc
    stage_acc<32, 16, LD64>(xacc, wm * 32, wn * 16, s.ba, false);
    __syncwarp();
#ifdef K1_DBG
    dump_bf(dbg, 36864, s.ba, 64*LD64);
    __syncthreads();
#endif

    // ---------------- w = M @ kb ; u = M @ vb ----------------
    float wacc[2][4][4], uacc[2][4][4];
    zero_acc<32, 32>(wacc);
    zero_acc<32, 32>(uacc);
    dot_mn<32, 32>(wacc, s.ba, wm * 32, LD64, s.kb, wn * 32, LD128, 64);
    dot_mn<32, 32>(uacc, s.ba, wm * 32, LD64, s.vb, wn * 32, LD128, 64);
    stage_acc<32, 32, LD128>(wacc, wm * 32, wn * 32, s.kb, false);
    stage_acc<32, 32, LD128>(uacc, wm * 32, wn * 32, s.vb, false);
    __syncthreads();
#ifdef K1_DBG
    dump_bf(dbg, 40960, s.kb, 64*LD128);
    dump_bf(dbg, 49152, s.vb, 64*LD128);
    __syncthreads();
#endif

#ifdef K1_EARLY
    if (K1_EARLY == 3) return;
#endif
    // ---------------- Aqk = tril(ql @ krT, incl diag) ; ql precomputed in prologue ----------------
    float aqk[2][2][4];
    zero_acc<32, 16>(aqk);
    dot_mn<32, 16>(aqk, s.qlq, wm * 32, LD128, s.krt, wn * 16, LD64, 128);
#pragma unroll
    for (int mi = 0; mi < 2; mi++)
#pragma unroll
        for (int ni = 0; ni < 2; ni++)
#pragma unroll
            for (int p = 0; p < 4; p++) {
                int r = wm * 32 + mi * 16 + (l >> 2) + (p >= 2 ? 8 : 0);
                int c = wn * 16 + ni * 8 + ((l & 3) * 2 + (p & 1));
                aqk[mi][ni][p] = (r >= c) ? aqk[mi][ni][p] : 0.f;   // ql already has scale
            }
    stage_acc<32, 16, LD64>(aqk, wm * 32, wn * 16, s.bb, false);
    __syncthreads();
#ifdef K1_DBG
    dump_bf(dbg, 57344, s.bb, 64*LD64);
#endif

    // ---------------- qe = ql*d - Aqk@w ; t = Aqk@u ----------------
    float qeacc[2][4][4], tacc[2][4][4];
    zero_acc<32, 32>(qeacc);
    zero_acc<32, 32>(tacc);
    dot_mn<32, 32>(qeacc, s.bb, wm * 32, LD64, s.kb, wn * 32, LD128, 64);
    dot_mn<32, 32>(tacc, s.bb, wm * 32, LD64, s.vb, wn * 32, LD128, 64);
    // pack qe with + ql[r][c]*d[c], stage into kl (element-wise read then overwrite)
#pragma unroll
    for (int mi = 0; mi < 2; mi++)
#pragma unroll
        for (int ni = 0; ni < 4; ni += 2) {
            int r = wm * 32 + mi * 16;
            int cc = wn * 32 + ni * 8;
            int lr = r + (l % 8) + ((l / 8) & 1) * 8;
            int lc = cc + (l / 16) * 8;
            int off = o128(lr, lc);
            int rr = r + (l >> 2);                 // fragment row (address row differs!)
            int cb = cc + 2 * (l & 3);             // fragment cols
            float q00 = bf2f(s.qlq[o128(rr, cb)]) * s.d_s[cb];
            float q01 = bf2f(s.qlq[o128(rr, cb + 1)]) * s.d_s[cb + 1];
            float q02 = bf2f(s.qlq[o128(rr + 8, cb)]) * s.d_s[cb];
            float q03 = bf2f(s.qlq[o128(rr + 8, cb + 1)]) * s.d_s[cb + 1];
            float q10 = bf2f(s.qlq[o128(rr, cb + 8)]) * s.d_s[cb + 8];
            float q11 = bf2f(s.qlq[o128(rr, cb + 9)]) * s.d_s[cb + 9];
            float q12 = bf2f(s.qlq[o128(rr + 8, cb + 8)]) * s.d_s[cb + 8];
            float q13 = bf2f(s.qlq[o128(rr + 8, cb + 9)]) * s.d_s[cb + 9];
            stm_x4(s.kl + off,
                   pack_bf16(q00 - qeacc[mi][ni][0], q01 - qeacc[mi][ni][1]),
                   pack_bf16(q02 - qeacc[mi][ni][2], q03 - qeacc[mi][ni][3]),
                   pack_bf16(q10 - qeacc[mi][ni + 1][0], q11 - qeacc[mi][ni + 1][1]),
                   pack_bf16(q12 - qeacc[mi][ni + 1][2], q13 - qeacc[mi][ni + 1][3]));
        }
    __syncwarp();
    __syncthreads();
    tile_to_gmem(s.kl, qe + outQ, 0, HK);
    __syncthreads();
    stage_acc<32, 32, LD128>(tacc, wm * 32, wn * 32, s.kl, false);
    __syncthreads();
    tile_to_gmem(s.kl, tt + outQ, 0, HK);
    __syncthreads();

#ifdef K1_EARLY
    if (K1_EARLY == 4) return;
#endif
    // ---------------- P = -(krt @ w) ; C = krt @ u  (all 8 warps, [64,32] tiles) --------
    // output rows 0-63 stage into qlq, rows 64-127 into kl (both free); kb/vb hold w/u.
    long pbase = ((i_b * H + i_h) * NT + pid_t) * 128 * 128;
    {
        float pacc[4][4][4];
        zero_acc<64, 32>(pacc);
        dot_mn<64, 32>(pacc, s.krt, wm * 64, LD64, s.kb, wn * 32, LD128, 64);
        __syncthreads();
        stage_acc<64, 32, LD128>(pacc, 0, wn * 32, wm ? s.kl : s.qlq, true);
        __syncthreads();
        tile_to_gmem(s.qlq, P + pbase, 0, 128);
        tile_to_gmem(s.kl, P + pbase + 64 * 128, 0, 128);
        __syncthreads();
    }
    {
        float cacc[4][4][4];
        zero_acc<64, 32>(cacc);
        dot_mn<64, 32>(cacc, s.krt, wm * 64, LD64, s.vb, wn * 32, LD128, 64);
        __syncthreads();
        stage_acc<64, 32, LD128>(cacc, 0, wn * 32, wm ? s.kl : s.qlq, false);
        __syncthreads();
        tile_to_gmem(s.qlq, Cm + pbase, 0, 128);
        tile_to_gmem(s.kl, Cm + pbase + 64 * 128, 0, 128);
        __syncthreads();
    }
}

void kda_k1_launch(torch::Tensor q, torch::Tensor k, torch::Tensor v, torch::Tensor g,
                   torch::Tensor beta, torch::Tensor P, torch::Tensor C, torch::Tensor Pd,
                   torch::Tensor qe, torch::Tensor t, long T, long H, long NT, double scale
#ifdef K1_DBG
                   , torch::Tensor dbg
#endif
                   ) {
    dim3 grid((unsigned)NT, (unsigned)(q.size(0) * H));
    size_t smem = sizeof(Smem);
    static bool set = false;
    if (!set) {
        cudaError_t e = cudaFuncSetAttribute(kda_k1, cudaFuncAttributeMaxDynamicSharedMemorySize, (int)smem);
        TORCH_CHECK(e == cudaSuccess, "setattr failed: ", cudaGetErrorString(e), " smem=", smem);
        set = true;
    }
    cudaError_t el = cudaGetLastError();
    TORCH_CHECK(el == cudaSuccess, "pre-launch: ", cudaGetErrorString(el));
    kda_k1<<<grid, 256, smem>>>(
        reinterpret_cast<const __nv_bfloat16*>(q.data_ptr()),
        reinterpret_cast<const __nv_bfloat16*>(k.data_ptr()),
        reinterpret_cast<const __nv_bfloat16*>(v.data_ptr()),
        g.data_ptr<float>(),
        reinterpret_cast<const __nv_bfloat16*>(beta.data_ptr()),
        reinterpret_cast<__nv_bfloat16*>(P.data_ptr()),
        reinterpret_cast<__nv_bfloat16*>(C.data_ptr()),
        Pd.data_ptr<float>(),
        reinterpret_cast<__nv_bfloat16*>(qe.data_ptr()),
        reinterpret_cast<__nv_bfloat16*>(t.data_ptr()),
        T, (int)H, NT, (float)scale
#ifdef K1_DBG
        , dbg.data_ptr<float>()
#endif
        );
}
"""


def _build_cuda_k1():
    """Compile the custom CUDA K1 kernel; returns None if unavailable."""
    try:
        if not torch.cuda.is_available():
            return None
        major, _ = torch.cuda.get_device_capability(0)
        if major != 12:
            return None
        from torch.utils.cpp_extension import load_inline
        bdir = os.path.join(os.path.dirname(os.path.abspath(__file__)), "_build")
        os.makedirs(bdir, exist_ok=True)
        return load_inline(
            name="kda_cuda_k1_v3",
            cpp_sources="void kda_k1_launch(torch::Tensor q, torch::Tensor k, torch::Tensor v, "
                        "torch::Tensor g, torch::Tensor beta, torch::Tensor P, torch::Tensor C, "
                        "torch::Tensor Pd, torch::Tensor qe, torch::Tensor t, long T, long H, "
                        "long NT, double scale);",
            cuda_sources=_CUDA_K1_SRC,
            functions=["kda_k1_launch"],
            extra_cuda_cflags=["-O3", "--use_fast_math", "-arch=sm_120a", "-std=c++20"],
            build_directory=bdir,
            verbose=False,
        )
    except Exception:
        return None


_CUDA_K1 = None
_CUDA_K1_TRIED = False


def _run_k1(q, k, v, g, beta, P_, C_, Pd_, qe_, t_, T, H, NT, scale, BH=1):
    """Dispatch K1 to the CUDA kernel when available, else Triton."""
    global _CUDA_K1, _CUDA_K1_TRIED
    if not _CUDA_K1_TRIED:
        _CUDA_K1 = _build_cuda_k1()
        _CUDA_K1_TRIED = True
    if _CUDA_K1 is not None:
        _CUDA_K1.kda_k1_launch(q, k, v, g, beta, P_, C_, Pd_, qe_, t_, T, H, NT, scale)
    else:
        _kda_k1[(NT, BH_global(q))](
            q, k, v, g, beta, P_, C_, Pd_, qe_, t_, T, H, NT, scale, num_warps=8)


_BUF_CACHE: dict = {}


class Model(nn.Module):
    """KDA forward (chunk form). No learned parameters; all inputs are activations."""

    def __init__(self, B: int, T: int, H: int, K: int, V: int, chunk_size: int = 64):
        super().__init__()
        self.B, self.T, self.H, self.K, self.V = B, T, H, K, V
        self.chunk_size = chunk_size
        self.scale = float(K) ** -0.5
        self.register_buffer("_dummy", torch.zeros(1), persistent=False)

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


# Module-level shape shims (overridden by check.py / benchmark.py per shape).
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]

20260822_053629_zai-claude_glm-5.3_02_kda_cutlass