"""Kimi Delta Attention (KDA) forward, chunk form -- custom Triton kernels for SM90. Two-kernel chunk-parallel implementation: Kernel 1 ("prep"), one program per chunk (parallel over B*H*NT): - in-chunk cumsum of the fp32 log-decay g (converted to base-2) - Aqk = tril( (q*scale*exp2(G)) @ (k*exp2(-G))^T ) - L = strict_tril( -beta_row * (k*exp2(G)) @ (k*exp2(-G))^T ) - X = (I - L)^-1 as the finite Neumann series sum_{j<64} L^j, built with repeated squaring: prod_{i=0..5} (I + L^{2^i}) (L^64 = 0) - w = X @ (beta * k * exp2(G)), u = X @ (beta * v) - stores qg = q*scale*exp2(G), kg = k*exp2(G_last - G), G_last Kernel 2 ("scan"), one program per (batch-head, value-block), sequential over chunks, state S kept in registers: v_n = u - w @ S o = qg @ S + Aqk @ v_n S = exp2(G_last) * S + kg^T @ v_n Matches the math of the reference chunk formulation (see reference.py). """ from __future__ import annotations import torch import torch.nn as nn import triton import triton.language as tl RCP_LN2 = tl.constexpr(1.4426950408889634) # --------------------------------------------------------------------------- # Kernel 1: per-chunk preparation (fully parallel) # --------------------------------------------------------------------------- @triton.jit(do_not_specialize=["T", "NT"]) def _kda_prep_kernel( q_ptr, k_ptr, g_ptr, beta_ptr, a_ptr, qg_ptr, kg_ptr, gl_ptr, l_ptr, f_ptr, scale, T, NT, H: tl.constexpr, K: tl.constexpr, BT: tl.constexpr, ): """Streaming phase: reads q,k,g,beta once; writes qg, kg, gl, Aqk, L, f. f = beta * k * exp2(G) so the compute kernel can form w = X @ f without touching q/k/g again. """ i_n = tl.program_id(0) i_bh = tl.program_id(1) i_b = i_bh // H i_h = i_bh % H BK: tl.constexpr = K // 2 # 64 for K=128 r64 = tl.arange(0, BT) tok = i_b * T + i_n * BT + r64 # (BT,) absolute tokens row_off = ((tok * H + i_h) * K) # (BT,) row offsets slab = (i_bh * NT + i_n) * BT cols0 = tl.arange(0, BK) cols1 = BK + tl.arange(0, BK) b_beta = tl.load(beta_ptr + tok * H + i_h).to(tl.float32) acc_a = tl.zeros([BT, BT], dtype=tl.float32) acc_l = tl.zeros([BT, BT], dtype=tl.float32) # ---- K slice 0 ---- b_g = tl.load(g_ptr + row_off[:, None] + cols0[None, :]) G = tl.cumsum(b_g, 0) * RCP_LN2 b_q = tl.load(q_ptr + row_off[:, None] + cols0[None, :]).to(tl.float32) b_k = tl.load(k_ptr + row_off[:, None] + cols0[None, :]).to(tl.float32) e_pos = tl.exp2(G) e_neg = tl.exp2(-G) qg0 = (b_q * e_pos * scale).to(tl.bfloat16) kp0 = (b_k * e_pos).to(tl.bfloat16) kn0 = (b_k * e_neg).to(tl.bfloat16) f0 = (b_beta[:, None] * b_k * e_pos).to(tl.bfloat16) tl.store(qg_ptr + (slab + r64[:, None]) * K + cols0[None, :], qg0) tl.store(f_ptr + (slab + r64[:, None]) * K + cols0[None, :], f0) gl0 = tl.sum(tl.where(r64[:, None] == BT - 1, G, 0.0), 0) egl0 = tl.exp2(gl0) tl.store(gl_ptr + (i_bh * NT + i_n) * K + cols0, gl0) tl.store(kg_ptr + (slab + r64[:, None]) * K + cols0[None, :], (kn0 * egl0[None, :]).to(tl.bfloat16)) p_t0 = ((tok[None, :] * H + i_h) * K) + cols0[:, None] b_gT = tl.load(g_ptr + p_t0) GT = tl.cumsum(b_gT, 1) * RCP_LN2 b_kT = tl.load(k_ptr + p_t0).to(tl.float32) knT = (b_kT * tl.exp2(-GT)).to(tl.bfloat16) acc_a += tl.dot(qg0, knT) acc_l += tl.dot(kp0, knT) # ---- K slice 1 ---- b_g = tl.load(g_ptr + row_off[:, None] + cols1[None, :]) G = tl.cumsum(b_g, 0) * RCP_LN2 b_q = tl.load(q_ptr + row_off[:, None] + cols1[None, :]).to(tl.float32) b_k = tl.load(k_ptr + row_off[:, None] + cols1[None, :]).to(tl.float32) e_pos = tl.exp2(G) e_neg = tl.exp2(-G) qg1 = (b_q * e_pos * scale).to(tl.bfloat16) kp1 = (b_k * e_pos).to(tl.bfloat16) kn1 = (b_k * e_neg).to(tl.bfloat16) f1 = (b_beta[:, None] * b_k * e_pos).to(tl.bfloat16) tl.store(qg_ptr + (slab + r64[:, None]) * K + cols1[None, :], qg1) tl.store(f_ptr + (slab + r64[:, None]) * K + cols1[None, :], f1) gl1 = tl.sum(tl.where(r64[:, None] == BT - 1, G, 0.0), 0) egl1 = tl.exp2(gl1) tl.store(gl_ptr + (i_bh * NT + i_n) * K + cols1, gl1) tl.store(kg_ptr + (slab + r64[:, None]) * K + cols1[None, :], (kn1 * egl1[None, :]).to(tl.bfloat16)) p_t1 = ((tok[None, :] * H + i_h) * K) + cols1[:, None] b_gT = tl.load(g_ptr + p_t1) GT = tl.cumsum(b_gT, 1) * RCP_LN2 b_kT = tl.load(k_ptr + p_t1).to(tl.float32) knT = (b_kT * tl.exp2(-GT)).to(tl.bfloat16) acc_a += tl.dot(qg1, knT) acc_l += tl.dot(kp1, knT) # ---- mask and store Aqk; store L = strict_tril(-beta_row * acc_l) ---- m_lo = r64[:, None] >= r64[None, :] tl.store(a_ptr + (slab + r64[:, None]) * BT + r64[None, :], tl.where(m_lo, acc_a, 0.0).to(a_ptr.dtype.element_ty)) L = tl.where(r64[:, None] > r64[None, :], -b_beta[:, None] * acc_l, 0.0) tl.store(l_ptr + (slab + r64[:, None]) * BT + r64[None, :], L.to(l_ptr.dtype.element_ty)) # --------------------------------------------------------------------------- # Kernel 1b: per-chunk triangular inverse + w/u (compute-bound) # --------------------------------------------------------------------------- @triton.jit(do_not_specialize=["T", "NT"]) def _kda_wu_kernel( v_ptr, beta_ptr, l_ptr, f_ptr, w_ptr, u_ptr, T, NT, H: tl.constexpr, K: tl.constexpr, V: tl.constexpr, BT: tl.constexpr, ): i_n = tl.program_id(0) i_bh = tl.program_id(1) i_b = i_bh // H i_h = i_bh % H BK: tl.constexpr = K // 2 r64 = tl.arange(0, BT) tok = i_b * T + i_n * BT + r64 slab = (i_bh * NT + i_n) * BT cols0 = tl.arange(0, BK) cols1 = BK + tl.arange(0, BK) b_beta = tl.load(beta_ptr + tok * H + i_h).to(tl.float32) L = tl.load(l_ptr + (slab + r64[:, None]) * BT + r64[None, :]).to(tl.float32) # X = prod_{i=0..5} (I + L^{2^i}) = sum_{j=0}^{63} L^j. Powers are built # on the fly so only a few 64x64 tiles are live at any time. I64 = (r64[:, None] == r64[None, :]).to(tl.float32) X = I64 + L P = L for _ in tl.static_range(5): P = tl.dot(P, P, input_precision="tf32") X = tl.dot(X, I64 + P, input_precision="tf32") Xb = X.to(tl.bfloat16) # w = X @ f b_f0 = tl.load(f_ptr + (slab + r64[:, None]) * K + cols0[None, :]) w0 = tl.dot(Xb, b_f0) tl.store(w_ptr + (slab + r64[:, None]) * K + cols0[None, :], w0.to(w_ptr.dtype.element_ty)) b_f1 = tl.load(f_ptr + (slab + r64[:, None]) * K + cols1[None, :]) w1 = tl.dot(Xb, b_f1) tl.store(w_ptr + (slab + r64[:, None]) * K + cols1[None, :], w1.to(w_ptr.dtype.element_ty)) # u = X @ (beta * v) b_v0 = tl.load(v_ptr + ((tok * H + i_h) * V)[:, None] + cols0[None, :]) u0 = tl.dot(Xb, (b_beta[:, None] * b_v0.to(tl.float32)).to(tl.bfloat16)) tl.store(u_ptr + (slab + r64[:, None]) * V + cols0[None, :], u0.to(u_ptr.dtype.element_ty)) b_v1 = tl.load(v_ptr + ((tok * H + i_h) * V)[:, None] + cols1[None, :]) u1 = tl.dot(Xb, (b_beta[:, None] * b_v1.to(tl.float32)).to(tl.bfloat16)) tl.store(u_ptr + (slab + r64[:, None]) * V + cols1[None, :], u1.to(u_ptr.dtype.element_ty)) # --------------------------------------------------------------------------- # Kernel 2: sequential inter-chunk scan (one program per head x value-block) # --------------------------------------------------------------------------- @triton.jit(do_not_specialize=["T", "NT"]) def _kda_scan_kernel( w_ptr, u_ptr, a_ptr, qg_ptr, kg_ptr, gl_ptr, o_ptr, T, NT, H: tl.constexpr, K: tl.constexpr, V: tl.constexpr, BT: tl.constexpr, BV: tl.constexpr, ): i_bh = tl.program_id(0) i_v = tl.program_id(1) i_b = i_bh // H i_h = i_bh % H r64 = tl.arange(0, BT) rk = tl.arange(0, K) colsv = i_v * BV + tl.arange(0, BV) S = tl.zeros([K, BV], dtype=tl.float32) for n in range(NT): slab = (i_bh * NT + n) * BT b_w = tl.load(w_ptr + (slab + r64[:, None]) * K + rk[None, :]) b_qg = tl.load(qg_ptr + (slab + r64[:, None]) * K + rk[None, :]) b_kg = tl.load(kg_ptr + (slab + r64[:, None]) * K + rk[None, :]) b_a = tl.load(a_ptr + (slab + r64[:, None]) * BT + r64[None, :]) b_u = tl.load(u_ptr + (slab + r64[:, None]) * V + colsv[None, :]).to(tl.float32) b_gl = tl.load(gl_ptr + (i_bh * NT + n) * K + rk) S_bf = S.to(tl.bfloat16) b_v = b_u - tl.dot(b_w, S_bf) b_v_bf = b_v.to(tl.bfloat16) b_o = tl.dot(b_qg, S_bf) + tl.dot(b_a, b_v_bf) tok = i_b * T + n * BT + r64 tl.store(o_ptr + ((tok * H + i_h) * V)[:, None] + colsv[None, :], b_o.to(o_ptr.dtype.element_ty)) S = S * tl.exp2(b_gl)[:, None] + tl.dot(tl.trans(b_kg), b_v_bf) # --------------------------------------------------------------------------- # Host wrapper # --------------------------------------------------------------------------- def _prep_config(): return {"num_warps": 4, "num_stages": 1} def _wu_config(): return {"num_warps": 4, "num_stages": 1} def _scan_config(): return {"BV": 32, "num_warps": 4, "num_stages": 3} 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 # No learned params; declare a dummy buffer so state_dict is well-defined. self.register_buffer("_dummy", torch.zeros(1), persistent=False) self._ws = None self._ws_key = None def _workspace(self, B, T, H, K, V, NT, device): key = (B, T, H, K, V) if self._ws_key != key: n_chunk_rows = B * H * NT * self.chunk_size cs = self.chunk_size self._ws = { "w": torch.empty(n_chunk_rows, K, device=device, dtype=torch.bfloat16), "u": torch.empty(n_chunk_rows, V, device=device, dtype=torch.bfloat16), "qg": torch.empty(n_chunk_rows, K, device=device, dtype=torch.bfloat16), "kg": torch.empty(n_chunk_rows, K, device=device, dtype=torch.bfloat16), "a": torch.empty(n_chunk_rows, cs, device=device, dtype=torch.bfloat16), "l": torch.empty(n_chunk_rows, cs, device=device, dtype=torch.bfloat16), "f": torch.empty(n_chunk_rows, K, device=device, dtype=torch.bfloat16), "gl": torch.empty(B * H * NT, K, device=device, dtype=torch.float32), } self._ws_key = key return self._ws def forward( self, q: torch.Tensor, k: torch.Tensor, v: torch.Tensor, g: torch.Tensor, beta: torch.Tensor, ) -> torch.Tensor: B, T, H, K = q.shape V = v.shape[-1] BT = self.chunk_size assert T % BT == 0 NT = T // BT ws = self._workspace(B, T, H, K, V, NT, q.device) o = torch.empty_like(v) cfg = _prep_config() _kda_prep_kernel[(NT, B * H)]( q, k, g, beta, ws["a"], ws["qg"], ws["kg"], ws["gl"], ws["l"], ws["f"], self.scale, T, NT, H=H, K=K, BT=BT, num_warps=cfg["num_warps"], num_stages=cfg["num_stages"], ) wcfg = _wu_config() _kda_wu_kernel[(NT, B * H)]( v, beta, ws["l"], ws["f"], ws["w"], ws["u"], T, NT, H=H, K=K, V=V, BT=BT, num_warps=wcfg["num_warps"], num_stages=wcfg["num_stages"], ) scfg = _scan_config() _kda_scan_kernel[(B * H, V // scfg["BV"])]( ws["w"], ws["u"], ws["a"], ws["qg"], ws["kg"], ws["gl"], o, T, NT, H=H, K=K, V=V, BT=BT, BV=scfg["BV"], num_warps=scfg["num_warps"], num_stages=scfg["num_stages"], ) return o # 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(): """Return a list of activations for one forward call. bf16 for q/k/v/beta; fp32 for the log-decay g (per FLA convention). """ 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 # log-decay: small negative numbers so exp(g) is in (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]