"""Kimi Delta Attention (KDA) chunk-parallel forward -- custom Triton kernels. Two-kernel design (both written from scratch; no library attention ops): K1 intra-chunk pass, grid (NT, B*H), one program per (batch*head, chunk): gc = cumsum(g) (chunk-local, per channel) P = k * exp(-gc) Kp = k * exp(gc) Qs = q*scale*exp(gc) G = Kp @ P^T (natural decay direction) L = diag(beta) tril(G, -1) Minv = (I + L)^-1 (blocked triangular inverse) u = Minv @ (beta*v) w = Minv @ (beta*Kp) Pt = P * exp(gc_last) d = exp(gc_last) Aqk = tril(Qs @ P^T) (j <= i) emits u, w, Qs, Pt (bf16), Aqk (bf16), d (fp32). K2 inter-chunk state pass, grid (NV, B*H), one program per (batch*head, V block): S = diag(d) S + Pt^T (u - w S) o = Qs S + Aqk (u - w S) The triangular inverse is factored as (I+L)^-1 = (I+N)^-1 Md with Md the inverse of the block-diagonal (BC x BC) part and N = Md @ L_off strictly block-lower (nilpotent). Md itself is evaluated by the exact product form (I+A)(I+A^2)(I+A^4)(I+A^8) with A = -L_diag and L_diag^BC = 0, which never forms large intermediate powers the way a plain Neumann series would. Precision: the (C,K) operand tiles are rounded to bf16 for the tensor-core MMAs (every accumulator stays fp32); the only values that need more than bf16 are the (C,C) inverse chain, which keeps tf32. That measured ~0.5% relative deviation from the all-tf32 version -- far inside the 0.05 tolerance -- and removed the ptxas spills that were costing ~1.6x on this kernel. Performance notes (RTX PRO 6000 / SM120): * the chain kernel's cost is its serial dependency (state -> dot -> state), so it wants the *smallest* useful V block: 16 columns with every (b,h) chain in flight measured 3x faster than 64-column blocks. * this harness flushes L2 (128 MB zero_) before each timed call, so the first touch of anything is a DRAM miss and every dirty flush line we evict costs a writeback: streamed inputs are loaded evict_first, the intermediate tiles that the chain kernel re-reads nv times are stored evict_last. """ from __future__ import annotations import torch import torch.nn as nn import triton import triton.language as tl _TF32 = tl.constexpr("tf32") # The harness flushes L2 (128 MB zero_) before every timed call, so the first # touch of anything is a DRAM miss and every dirty flush line we evict costs a # writeback. q/k/v/g are streamed once per program -> evict-first; the tiles # the chain kernel reads back nv times -> evict-last so they stay resident. _EV_FIRST = tl.constexpr("evict_first") _EV_LAST = tl.constexpr("evict_last") @triton.jit def _kda_intra_kernel( q_ptr, k_ptr, v_ptr, g_ptr, beta_ptr, u_ptr, w_ptr, qs_ptr, pt_ptr, aqk_ptr, d_ptr, scale, T, NT, H: tl.constexpr, K: tl.constexpr, V: tl.constexpr, C: tl.constexpr, BC: tl.constexpr, ): pid_c = tl.program_id(0) pid_bh = tl.program_id(1) b = pid_bh // H h = pid_bh - b * H rc = tl.arange(0, C) rk = tl.arange(0, K) rv = tl.arange(0, V) rows = pid_c * C + rc qk_off = (b * T + rows)[:, None] * (H * K) + h * K + rk[None, :] q = tl.load(q_ptr + qk_off, eviction_policy=_EV_FIRST) k = tl.load(k_ptr + qk_off, eviction_policy=_EV_FIRST) g = tl.load(g_ptr + qk_off, eviction_policy=_EV_FIRST) v = tl.load(v_ptr + (b * T + rows)[:, None] * (H * V) + h * V + rv[None, :], eviction_policy=_EV_FIRST) beta = tl.load(beta_ptr + (b * T + rows) * H + h) gc = tl.cumsum(g, axis=0) glast = tl.sum(tl.where(rc[:, None] == C - 1, gc, 0.0), axis=0) # (K,) # Keep the (C,K) operand tiles in bf16: they only ever feed tensor-core MMAs # whose accumulators are fp32, and the fp32 copies were the register hogs # (ptxas was spilling 24-70 words/thread, which cost more than the rounding). kf = k.to(tl.float32) em = tl.exp(-gc) eg = tl.exp(gc) P = (kf * em).to(tl.bfloat16) # k e^{-gc} Kp = (kf * eg).to(tl.bfloat16) # k e^{gc} Qs = (q.to(tl.float32) * (scale * eg)).to(tl.bfloat16) # q*scale e^{gc} P32 = P.to(tl.float32) # ---- intra-chunk Gram in the decaying direction: G[a,b] = k_a k_b e^{gc_a-gc_b} G = tl.dot(Kp, tl.trans(P), out_dtype=tl.float32) strict_lower = rc[:, None] > rc[None, :] L = tl.where(strict_lower, G, 0.0) * beta[:, None] # diag(beta) tril(G,-1) A = -L eye = (rc[:, None] == rc[None, :]).to(tl.float32) if C > BC: blk = (rc[:, None] // BC) == (rc[None, :] // BC) Ad = tl.where(blk, A, 0.0) else: Ad = A # Md = (I - Ad)^{-1} = (I+Ad)(I+Ad^2)(I+Ad^4)(I+Ad^8); Ad^BC = 0 for BC=16 M = eye + Ad A2 = tl.dot(Ad, Ad, input_precision=_TF32) M = M + tl.dot(M, A2, input_precision=_TF32) A4 = tl.dot(A2, A2, input_precision=_TF32) M = M + tl.dot(M, A4, input_precision=_TF32) A8 = tl.dot(A4, A4, input_precision=_TF32) Md = M + tl.dot(M, A8, input_precision=_TF32) if C > BC: Loff = L - tl.where(blk, L, 0.0) N = tl.dot(Md, Loff, input_precision=_TF32) # strictly block lower if C // BC > 2: N2 = tl.dot(N, N, input_precision=_TF32) N3 = tl.dot(N2, N, input_precision=_TF32) Q = eye - N + N2 - N3 else: Q = eye - N Minv = tl.dot(Q, Md, input_precision=_TF32) else: Minv = Md Mib = Minv.to(tl.bfloat16) betab = beta.to(tl.bfloat16) u = tl.dot(Mib, betab[:, None] * v, out_dtype=tl.float32) # (C,V) w = tl.dot(Mib, betab[:, None] * Kp, out_dtype=tl.float32) # (C,K) dl = tl.exp(glast) Pt = P32 * dl[None, :] # (C,K) Aqk = tl.dot(Qs, tl.trans(P), out_dtype=tl.float32) Aqk = tl.where(rc[:, None] >= rc[None, :], Aqk, 0.0) # keep j <= i (diag incl.) idx = pid_bh * NT + pid_c cu = u_ptr + idx * (C * V) + rc[:, None] * V + rv[None, :] ck = rc[:, None] * K + rk[None, :] tl.store(cu, u.to(tl.bfloat16), eviction_policy=_EV_LAST) tl.store(w_ptr + idx * (C * K) + ck, w.to(tl.bfloat16), eviction_policy=_EV_LAST) tl.store(qs_ptr + idx * (C * K) + ck, Qs.to(tl.bfloat16), eviction_policy=_EV_LAST) tl.store(pt_ptr + idx * (K * C) + rk[:, None] * C + rc[None, :], tl.trans(Pt).to(tl.bfloat16), eviction_policy=_EV_LAST) tl.store(aqk_ptr + idx * (C * C) + rc[:, None] * C + rc[None, :], Aqk.to(tl.bfloat16), eviction_policy=_EV_LAST) tl.store(d_ptr + idx * K + rk, dl, eviction_policy=_EV_LAST) @triton.jit def _kda_chain_kernel( u_ptr, w_ptr, qs_ptr, pt_ptr, aqk_ptr, d_ptr, o_ptr, T, NT, H: tl.constexpr, K: tl.constexpr, V: tl.constexpr, C: 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 rk = tl.arange(0, K) rc = tl.arange(0, C) rv = pid_v * BV + tl.arange(0, BV) S = tl.zeros((K, BV), dtype=tl.float32) for c in range(0, NT): idx = pid_bh * NT + c u = tl.load(u_ptr + idx * (C * V) + rc[:, None] * V + rv[None, :], eviction_policy=_EV_LAST).to(tl.float32) w = tl.load(w_ptr + idx * (C * K) + rc[:, None] * K + rk[None, :], eviction_policy=_EV_LAST) qs = tl.load(qs_ptr + idx * (C * K) + rc[:, None] * K + rk[None, :], eviction_policy=_EV_LAST) pt = tl.load(pt_ptr + idx * (K * C) + rk[:, None] * C + rc[None, :], eviction_policy=_EV_LAST) aqk = tl.load(aqk_ptr + idx * (C * C) + rc[:, None] * C + rc[None, :], eviction_policy=_EV_LAST) vi = u - tl.dot(w, S.to(tl.bfloat16)) # (C,BV) vib = vi.to(tl.bfloat16) o = tl.dot(qs, S.to(tl.bfloat16)) + tl.dot(aqk, vib) # (C,BV) d = tl.load(d_ptr + idx * K + rk, eviction_policy=_EV_LAST) S = S * d[:, None] + tl.dot(pt, vib) tl.store(o_ptr + (b * T + c * C + rc)[:, None] * (H * V) + h * V + rv[None, :], o.to(tl.bfloat16), eviction_policy=_EV_FIRST) # --------------------------------------------------------------------------- # host side # --------------------------------------------------------------------------- _DEF_C = 32 _DEF_BC = 16 def _pick_cfg(B, T, H, K, V): """(chunk, BV) for the intra kernel and BV for the chain kernel. The chain kernel is latency-bound on its serial dependency (state -> dot -> state), so its speed is set by how many independent (batch*head, V-block) chains can run concurrently. Measured on this part: 16-column V blocks with every (b,h) chain in flight is 3x faster than 64-column blocks, even though it re-reads the V-independent tiles more often (those stay L2-resident). """ C = 32 if T % 32 == 0 else 64 nv = 1 while nv < 8 and V // (nv * 2) >= 16: nv *= 2 return C, V // nv class Model(nn.Module): 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) self._bufs = None self._key = None def _get_bufs(self, B, T, H, K, V, C, BV, device): NT = T // C bhnt = B * H * NT key = (B, T, H, K, V, C, BV, str(device)) if self._key == key: return self._bufs f32 = torch.float32 bf = torch.bfloat16 bufs = { "u": torch.empty((bhnt, C, V), dtype=bf, device=device), "w": torch.empty((bhnt, C, K), dtype=bf, device=device), "qs": torch.empty((bhnt, C, K), dtype=bf, device=device), "pt": torch.empty((bhnt, K, C), dtype=bf, device=device), "aqk": torch.empty((bhnt, C, C), dtype=bf, device=device), "d": torch.empty((bhnt, K), dtype=f32, device=device), } self._bufs = bufs self._key = key return bufs def forward(self, q, k, v, g, beta): B, T, H, K = q.shape V = v.shape[-1] C, BV = _pick_cfg(B, T, H, K, V) if T % C or V % BV: C, BV = self.chunk_size, V NT = T // C dev = q.device q = q.contiguous() k = k.contiguous() v = v.contiguous() g = g.contiguous() beta = beta.contiguous() if q.dtype != torch.bfloat16: q = q.to(torch.bfloat16) k = k.to(torch.bfloat16) v = v.to(torch.bfloat16) beta = beta.to(torch.bfloat16) if g.dtype != torch.float32: g = g.to(torch.float32) bufs = self._get_bufs(B, T, H, K, V, C, BV, dev) BC = _DEF_BC if C >= _DEF_BC else C o = torch.empty((B, T, H, V), dtype=torch.bfloat16, device=dev) _kda_intra_kernel[(NT, B * H)]( q, k, v, g, beta, bufs["u"], bufs["w"], bufs["qs"], bufs["pt"], bufs["aqk"], bufs["d"], self.scale, T, NT, H=H, K=K, V=V, C=C, BC=BC, num_warps=4, num_stages=1, ) nv = V // BV _kda_chain_kernel[(nv, B * H)]( bufs["u"], bufs["w"], bufs["qs"], bufs["pt"], bufs["aqk"], bufs["d"], o, T, NT, H=H, K=K, V=V, C=C, BV=BV, num_warps=4, num_stages=3, ) return o def get_init_inputs(): return [2, 1024, 8, 128, 128, 64] def get_inputs(): torch.manual_seed(0) B, T, H, K, V = 2, 1024, 8, 128, 128 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]