"""Kimi Delta Attention (KDA) forward, chunk form — custom Triton kernels for B200 (SM100). Written from scratch against the chunk-parallel formulation in reference.py. The inter-chunk recurrence is linear in the K x V state S: S_{n+1} = diag(exp(gc_n[-1])) @ S_n + kg_n^T @ (u_n - w_n @ S_n) = An @ S_n + bn, An = diag(dlast_n) - kg_n^T @ w_n (K x K) bn = kg_n^T @ u_n (K x V) o_n = qg_n @ S_n + Aqk_n @ (u_n - w_n @ S_n) = (qg_n - Aqk_n @ w_n) @ S_n + Aqk_n @ u_n = qgp_n @ S_n + ou_n so everything except one short scan is embarrassingly parallel over chunks: k_pre chunk-local cumsum of g + every decay factor, derived from one exp tile + one reciprocal via a per-channel max anchor. k_A decay-weighted grams kp @ kn^T and qp @ kn^T (per-channel max anchor cancels), then the delta-rule WY transform (I - tril(beta k k^T))^{-1} via the nilpotent doubling identity (I-M)^{-1} = prod (I + M^{2^k}). k_wu w = A @ (exp(gc) k), u = A @ v. k_trans An (dense, diagonal folded) and bn per chunk. k_op qgp = qg - Aqk @ w, ou = Aqk @ u. k_compose two parallel tree levels build 4-chunk composed operators A4 = A3 A2 A1 A0, b4 = A3(A2(A1 b0 + b1) + b2) + b3. k_serial the only sequential kernel: NT/4 steps of S = A4 @ S + b4 (one dependent tensor-core dot per step), bf16 group snapshots. k_output parallel per group: replay the <=3 intra-group hops with An/bn and emit o = qgp @ S_n + ou for the 4 chunks. Model.forward wraps the pipeline in a per-shape CUDA graph: inputs are copied into static buffers on EVERY call and the graph replays the kernel pipeline, so each call fully recomputes from the live input values (no result caching). The graph only removes CPU launch overhead. KDA_DISABLE_GRAPH=1 runs eagerly. """ from __future__ import annotations import os import torch import torch.nn as nn import triton import triton.language as tl OP_TYPE = "linear_attention" SUPPORTED_PRECISIONS = ["bf16"] HARDWARE_REQUIRED = ["RTX_PRO_6000", "H100", "B200"] # --------------------------------------------------------------------------- # k_pre: chunk-local cumsum + all decay factors # --------------------------------------------------------------------------- @triton.jit def _kda_pre( q_ptr, k_ptr, g_ptr, kp_ptr, kn_ptr, qp_ptr, em_ptr, elm_ptr, dlast_ptr, H, T, NT, scale, BT: tl.constexpr, K: tl.constexpr, ): pid_n = tl.program_id(0) pid_bh = tl.program_id(1) b = pid_bh // H h = pid_bh % H offs_c = tl.arange(0, BT) offs_k = tl.arange(0, K) row = (b * T + pid_n * BT + offs_c) * H + h ptrs = row[:, None] * K + offs_k[None, :] g = tl.load(g_ptr + ptrs) # (BT, K) fp32 gc = tl.cumsum(g, axis=0) # chunk-local cumsum gl = tl.sum(g, axis=0) # last cumsum row, exact m = tl.max(gc, axis=0) # (K,) per-channel anchor ep = tl.exp(gc - m[None, :]) # <= 1 en = 1.0 / ep # exp(m - gc) em = tl.exp(m) egl = tl.exp(gl) elm = egl / em # exp(gl - m) kf = tl.load(k_ptr + ptrs).to(tl.float32) qf = tl.load(q_ptr + ptrs).to(tl.float32) * scale # Downstream kernels rebuild the remaining decay products per channel: # gk = kp * em, qg = qp * em, kg = kn * elm. tl.store(kp_ptr + ptrs, (kf * ep).to(tl.bfloat16)) tl.store(kn_ptr + ptrs, (kf * en).to(tl.bfloat16)) tl.store(qp_ptr + ptrs, (qf * ep).to(tl.bfloat16)) kbase = (pid_bh * NT + pid_n) * K + offs_k tl.store(em_ptr + kbase, em) tl.store(elm_ptr + kbase, elm) tl.store(dlast_ptr + kbase, egl) # --------------------------------------------------------------------------- # k_A: triangular inverse (WY transform); axis2 program 1 computes Aqk instead # --------------------------------------------------------------------------- @triton.jit def _kda_A( kp_ptr, kn_ptr, qp_ptr, beta_ptr, a_ptr, aqk_ptr, H, T, NT, BT: tl.constexpr, K: tl.constexpr, ): pid_n = tl.program_id(0) pid_bh = tl.program_id(1) which = tl.program_id(2) b = pid_bh // H h = pid_bh % H offs_c = tl.arange(0, BT) offs_k = tl.arange(0, K) row = (b * T + pid_n * BT + offs_c) * H + h ptrs = row[:, None] * K + offs_k[None, :] kn = tl.load(kn_ptr + ptrs) knt = tl.trans(kn) base = (pid_bh * NT + pid_n) * BT * BT idx = base + offs_c[:, None] * BT + offs_c[None, :] if which == 1: qp = tl.load(qp_ptr + ptrs) aqk = tl.dot(qp, knt) # (BT, BT) fp32 aqk = tl.where(offs_c[:, None] >= offs_c[None, :], aqk, 0.0) tl.store(aqk_ptr + idx, aqk.to(tl.bfloat16)) else: kp = tl.load(kp_ptr + ptrs) # (BT, K) bf16 afull = tl.dot(kp, knt) # (BT, BT) fp32 beta = tl.load(beta_ptr + row).to(tl.float32) # (BT,) strict_lower = offs_c[:, None] > offs_c[None, :] mm = tl.where(strict_lower, -(afull * beta[:, None]), 0.0) # (I-M)^{-1} = (I+M)(I+M^2)(I+M^4)(I+M^8)(I+M^16)(I+M^32); M^64 = 0. eye = tl.where(offs_c[:, None] == offs_c[None, :], 1.0, 0.0) x = eye + mm p = tl.dot(mm, mm) for _ in range(4): x = x + tl.dot(x, p) p = tl.dot(p, p) x = x + tl.dot(x, p) ab = (x * beta[None, :]).to(tl.bfloat16) # column scale by beta tl.store(a_ptr + idx, ab) # --------------------------------------------------------------------------- # k_wu: axis2 program 0: w = A @ gk; program 1: u = A @ v # --------------------------------------------------------------------------- @triton.jit def _kda_wu( a_ptr, kp_ptr, em_ptr, v_ptr, w_ptr, u_ptr, H, T, NT, BT: tl.constexpr, K: tl.constexpr, V: tl.constexpr, ): pid_n = tl.program_id(0) pid_bh = tl.program_id(1) which = tl.program_id(2) b = pid_bh // H h = pid_bh % H offs_c = tl.arange(0, BT) offs_k = tl.arange(0, K) row = (b * T + pid_n * BT + offs_c) * H + h a_base = (pid_bh * NT + pid_n) * BT * BT a = tl.load(a_ptr + a_base + offs_c[:, None] * BT + offs_c[None, :]) ptr = row[:, None] * K + offs_k[None, :] if which == 0: kp = tl.load(kp_ptr + ptr).to(tl.float32) em = tl.load(em_ptr + (pid_bh * NT + pid_n) * K + offs_k) gk = (kp * em[None, :]).to(tl.bfloat16) # k * exp(gc) w = tl.dot(a, gk) tl.store(w_ptr + ptr, w.to(tl.bfloat16)) else: v = tl.load(v_ptr + ptr) u = tl.dot(a, v) tl.store(u_ptr + ptr, u.to(tl.bfloat16)) # --------------------------------------------------------------------------- # k_trans: An = diag(dlast) - kg^T @ w (dense), bn = kg^T @ u # grid axis2: (K row block) x (an | bn) # --------------------------------------------------------------------------- @triton.jit def _kda_trans( w_ptr, u_ptr, kn_ptr, elm_ptr, dlast_ptr, an_ptr, bn_ptr, H, T, NT, BT: tl.constexpr, K: tl.constexpr, V: tl.constexpr, BKR: tl.constexpr, ): pid_n = tl.program_id(0) pid_bh = tl.program_id(1) pid_z = tl.program_id(2) pid_kb = pid_z % (K // BKR) which = pid_z // (K // BKR) b = pid_bh // H h = pid_bh % H offs_c = tl.arange(0, BT) offs_k = tl.arange(0, K) offs_kr = pid_kb * BKR + tl.arange(0, BKR) row = (b * T + pid_n * BT + offs_c) * H + h knr = tl.load(kn_ptr + row[:, None] * K + offs_kr[None, :]).to(tl.float32) elmr = tl.load(elm_ptr + (pid_bh * NT + pid_n) * K + offs_kr) kgt = (knr * elmr[None, :]).to(tl.bfloat16) # k * exp(gl - gc), (BT, BKR) base = (pid_bh * NT + pid_n) * K + pid_kb * BKR if which == 0: w = tl.load(w_ptr + row[:, None] * K + offs_k[None, :]) # (BT, K) akw = tl.dot(tl.trans(kgt), w) # (BKR, K) dlr = tl.load(dlast_ptr + (pid_bh * NT + pid_n) * K + offs_kr) an = tl.where(offs_kr[:, None] == offs_k[None, :], dlr[:, None] - akw, -akw) tl.store(an_ptr + base * K + tl.arange(0, BKR)[:, None] * K + offs_k[None, :], an.to(tl.bfloat16)) else: u = tl.load(u_ptr + row[:, None] * V + offs_k[None, :]) # (BT, V) bn = tl.dot(tl.trans(kgt), u) # (BKR, V) tl.store(bn_ptr + base * V + tl.arange(0, BKR)[:, None] * V + offs_k[None, :], bn.to(tl.bfloat16)) # --------------------------------------------------------------------------- # k_op: axis2 program 0: qgp = qg - Aqk @ w; program 1: ou = Aqk @ u # --------------------------------------------------------------------------- @triton.jit def _kda_op( w_ptr, u_ptr, qp_ptr, em_ptr, aqk_ptr, qgp_ptr, ou_ptr, H, T, NT, BT: tl.constexpr, K: tl.constexpr, V: tl.constexpr, ): pid_n = tl.program_id(0) pid_bh = tl.program_id(1) which = tl.program_id(2) b = pid_bh // H h = pid_bh % H offs_c = tl.arange(0, BT) offs_k = tl.arange(0, K) row = (b * T + pid_n * BT + offs_c) * H + h ptr = row[:, None] * K + offs_k[None, :] base = (pid_bh * NT + pid_n) * BT * BT aqk = tl.load(aqk_ptr + base + offs_c[:, None] * BT + offs_c[None, :]) if which == 0: w = tl.load(w_ptr + ptr) qp = tl.load(qp_ptr + ptr).to(tl.float32) em = tl.load(em_ptr + (pid_bh * NT + pid_n) * K + offs_k) qg = qp * em[None, :] # q * scale * exp(gc) qgp = qg - tl.dot(aqk, w) tl.store(qgp_ptr + ptr, qgp.to(tl.bfloat16)) else: u = tl.load(u_ptr + ptr) ou = tl.dot(aqk, u) tl.store(ou_ptr + ptr, ou.to(tl.bfloat16)) # --------------------------------------------------------------------------- # k_compose: one tree level; group pairs (lo, hi) -> hi ∘ lo # Aout = Ahi @ Alo, bout = Ahi @ blo + bhi. Column-split for registers. # --------------------------------------------------------------------------- @triton.jit def _kda_compose( alo_ptr, blo_ptr, ahi_ptr, bhi_ptr, aout_ptr, bout_ptr, NPAIR_PER_BH, K: tl.constexpr, V: tl.constexpr, BC: tl.constexpr, ): pid_p = tl.program_id(0) # pair index within bh pid_bh = tl.program_id(1) pid_c = tl.program_id(2) # column block offs_k = tl.arange(0, K) offs_c = pid_c * BC + tl.arange(0, BC) lo = (pid_bh * 2 * NPAIR_PER_BH + 2 * pid_p) hi = lo + 1 out = pid_bh * NPAIR_PER_BH + pid_p ahi = tl.load(ahi_ptr + hi * K * K + offs_k[:, None] * K + offs_k[None, :]) alo_c = tl.load(alo_ptr + lo * K * K + offs_k[:, None] * K + offs_c[None, :]) aout = tl.dot(ahi, alo_c) tl.store(aout_ptr + out * K * K + offs_k[:, None] * K + offs_c[None, :], aout.to(tl.bfloat16)) blo_c = tl.load(blo_ptr + lo * K * V + offs_k[:, None] * V + offs_c[None, :]) bhi_c = tl.load(bhi_ptr + hi * K * V + offs_k[:, None] * V + offs_c[None, :]).to(tl.float32) bout = tl.dot(ahi, blo_c) + bhi_c tl.store(bout_ptr + out * K * V + offs_k[:, None] * V + offs_c[None, :], bout.to(tl.bfloat16)) # --------------------------------------------------------------------------- # k_serial: S <- A4 @ S + b4 over NT/4 groups; bf16 snapshots per group # --------------------------------------------------------------------------- @triton.jit def _kda_serial( a4_ptr, b4_ptr, snap_ptr, NG, K: tl.constexpr, V: tl.constexpr, BV: tl.constexpr, NSTAGE: tl.constexpr, ): pid_bh = tl.program_id(0) pid_v = tl.program_id(1) offs_k = tl.arange(0, K) offs_v = pid_v * BV + tl.arange(0, BV) s = tl.zeros((K, BV), dtype=tl.float32) for n in tl.range(NG, num_stages=NSTAGE): base = pid_bh * NG + n sb = s.to(tl.bfloat16) tl.store(snap_ptr + base * K * V + offs_k[:, None] * V + offs_v[None, :], sb) a4 = tl.load(a4_ptr + base * K * K + offs_k[:, None] * K + offs_k[None, :]) b4 = tl.load(b4_ptr + base * K * V + offs_k[:, None] * V + offs_v[None, :]).to(tl.float32) s = b4 + tl.dot(a4, sb) # --------------------------------------------------------------------------- # k_output: per 4-chunk group, replay intra-group hops and emit o # --------------------------------------------------------------------------- @triton.jit def _kda_output( qgp_ptr, ou_ptr, an_ptr, bn_ptr, snap_ptr, o_ptr, H, T, NT, NG, BT: tl.constexpr, K: tl.constexpr, V: tl.constexpr, BV: tl.constexpr, L: tl.constexpr, ): pid_g = tl.program_id(0) # group index within bh pid_bhv = tl.program_id(1) NV: tl.constexpr = V // BV pid_bh = pid_bhv // NV pid_v = pid_bhv % NV b = pid_bh // H h = pid_bh % H offs_c = tl.arange(0, BT) offs_k = tl.arange(0, K) offs_v = pid_v * BV + tl.arange(0, BV) sbase = pid_bh * NG + pid_g sb = tl.load(snap_ptr + sbase * K * V + offs_k[:, None] * V + offs_v[None, :]) for j in tl.static_range(L): n = pid_g * L + j row = (b * T + n * BT + offs_c) * H + h qgp = tl.load(qgp_ptr + row[:, None] * K + offs_k[None, :]) ou = tl.load(ou_ptr + row[:, None] * V + offs_v[None, :]).to(tl.float32) o = tl.dot(qgp, sb) + ou tl.store(o_ptr + row[:, None] * V + offs_v[None, :], o.to(tl.bfloat16)) if j < L - 1: nbase = pid_bh * NT + n an = tl.load(an_ptr + nbase * K * K + offs_k[:, None] * K + offs_k[None, :]) bn = tl.load(bn_ptr + nbase * K * V + offs_k[:, None] * V + offs_v[None, :]).to(tl.float32) sb = (bn + tl.dot(an, sb)).to(tl.bfloat16) # --------------------------------------------------------------------------- # Python driver # --------------------------------------------------------------------------- _BT = 64 _BV_SER = int(os.environ.get("KDA_BV_SER", "16")) _NSTAGE = 3 _BV_OUT = 64 class _ShapeState: """Per-shape static workspace + optional CUDA graph.""" def __init__(self, B, T, H, K, V, scale, device): self.B, self.T, self.H, self.K, self.V = B, T, H, K, V self.scale = scale NT = T // _BT # Composed-group size: the serial scan runs NT/L steps; the output pass # replays L-1 intra-group hops in parallel. Deeper composition only # pays off when the scan is long. L = int(os.environ.get("KDA_L", "0")) or (4 if NT >= 64 else 2) NG = NT // L self.NT, self.NG, self.L = NT, NG, L BH = B * H dt = torch.bfloat16 dev = device e = torch.empty # q/k/v/beta live in one flat bf16 buffer so one cat() refreshes them. nq = B * T * H * K nv = B * T * H * V nb = B * T * H self.flat = e(nq + nq + nv + nb, dtype=dt, device=dev) self.q = self.flat[:nq].view(B, T, H, K) self.k = self.flat[nq:2 * nq].view(B, T, H, K) self.v = self.flat[2 * nq:2 * nq + nv].view(B, T, H, V) self.beta = self.flat[2 * nq + nv:].view(B, T, H) self.g = e(B, T, H, K, dtype=torch.float32, device=dev) self.kp = e(B, T, H, K, dtype=dt, device=dev) self.kn = e(B, T, H, K, dtype=dt, device=dev) self.qp = e(B, T, H, K, dtype=dt, device=dev) self.em = e(BH, NT, K, dtype=torch.float32, device=dev) self.elm = e(BH, NT, K, dtype=torch.float32, device=dev) self.dlast = e(BH, NT, K, dtype=torch.float32, device=dev) self.a = e(BH, NT, _BT, _BT, dtype=dt, device=dev) self.aqk = e(BH, NT, _BT, _BT, dtype=dt, device=dev) self.w = e(B, T, H, K, dtype=dt, device=dev) self.u = e(B, T, H, V, dtype=dt, device=dev) self.an = e(BH, NT, K, K, dtype=dt, device=dev) self.bn = e(BH, NT, K, V, dtype=dt, device=dev) self.a2 = e(BH, NT // 2, K, K, dtype=dt, device=dev) self.b2 = e(BH, NT // 2, K, V, dtype=dt, device=dev) if L == 4: self.a4 = e(BH, NG, K, K, dtype=dt, device=dev) self.b4 = e(BH, NG, K, V, dtype=dt, device=dev) else: self.a4, self.b4 = self.a2, self.b2 self.qgp = e(B, T, H, K, dtype=dt, device=dev) self.ou = e(B, T, H, V, dtype=dt, device=dev) self.snap = e(BH, NG, K, V, dtype=dt, device=dev) self.o = e(B, T, H, V, dtype=dt, device=dev) self.graph = None def launch(self): B, T, H, K, V, NT, NG = self.B, self.T, self.H, self.K, self.V, self.NT, self.NG BH = B * H _kda_pre[(NT, BH)]( self.q, self.k, self.g, self.kp, self.kn, self.qp, self.em, self.elm, self.dlast, H, T, NT, self.scale, BT=_BT, K=K, num_warps=4) _kda_A[(NT, BH, 2)]( self.kp, self.kn, self.qp, self.beta, self.a, self.aqk, H, T, NT, BT=_BT, K=K, num_warps=4) _kda_wu[(NT, BH, 2)]( self.a, self.kp, self.em, self.v, self.w, self.u, H, T, NT, BT=_BT, K=K, V=V, num_warps=4) _kda_op[(NT, BH, 2)]( self.w, self.u, self.qp, self.em, self.aqk, self.qgp, self.ou, H, T, NT, BT=_BT, K=K, V=V, num_warps=4) _kda_trans[(NT, BH, 4)]( self.w, self.u, self.kn, self.elm, self.dlast, self.an, self.bn, H, T, NT, BT=_BT, K=K, V=V, BKR=K // 2, num_warps=4) _kda_compose[(NT // 2, BH, 2)]( self.an, self.bn, self.an, self.bn, self.a2, self.b2, NT // 2, K=K, V=V, BC=K // 2, num_warps=8) if self.L == 4: _kda_compose[(NG, BH, 2)]( self.a2, self.b2, self.a2, self.b2, self.a4, self.b4, NG, K=K, V=V, BC=K // 2, num_warps=8) _kda_serial[(BH, V // _BV_SER)]( self.a4, self.b4, self.snap, NG, K=K, V=V, BV=_BV_SER, NSTAGE=_NSTAGE, num_warps=4) _kda_output[(NG, BH * (V // _BV_OUT))]( self.qgp, self.ou, self.an, self.bn, self.snap, self.o, H, T, NT, NG, BT=_BT, K=K, V=V, BV=_BV_OUT, L=self.L, num_warps=4) def run(self, q, k, v, g, beta, use_graph): # Fresh input values are copied into the static buffers on EVERY call; # the kernels then recompute the output from those values. self.g.copy_(g) torch.cat( [q.reshape(-1), k.reshape(-1), v.reshape(-1), beta.reshape(-1)], out=self.flat, ) if not use_graph: self.launch() return self.o if self.graph is None: torch.cuda.synchronize() for _ in range(3): self.launch() # Triton JIT + allocator warmup torch.cuda.synchronize() gr = torch.cuda.CUDAGraph() with torch.cuda.graph(gr): self.launch() self.graph = gr self.graph.replay() return self.o _STATES: dict = {} def kda_chunk_forward(q, k, v, g, beta, scale, chunk_size=64): B, T, H, K = q.shape V = v.shape[-1] assert chunk_size == _BT and T % (_BT * 2) == 0 and K == V key = (B, T, H, K, V, q.device.index) st = _STATES.get(key) if st is None: st = _ShapeState(B, T, H, K, V, scale, q.device) _STATES[key] = st use_graph = os.environ.get("KDA_DISABLE_GRAPH", "0") != "1" return st.run(q, k, v, g, beta, use_graph) 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_chunk_forward(q, k, v, g, beta, scale=self.scale, chunk_size=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]