"""Kimi Delta Attention chunk-form forward for SM120. Custom Triton kernels: 1. intra-chunk Akk (per-channel decay) 2. 16x16-block unit-lower inverse times beta 3. WY factors w, u plus Aqk / qg / kbar 4. tiled inter-chunk state recurrence and output CUDA graphs replay the four launches on the hot path. """ from __future__ import annotations 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"] @triton.jit def _build_A_kernel( k_ptr, g_ptr, beta_ptr, A_ptr, T, H, NT, BT: tl.constexpr, BK: tl.constexpr, ): i_t = tl.program_id(0) i_bh = tl.program_id(1) i_b = i_bh // H i_h = i_bh % H t0 = i_t * BT K = 128 o_c = tl.arange(0, BT) o_k = tl.arange(0, BK) b_A = tl.zeros([BT, BT], dtype=tl.float32) b_k = tl.load(k_ptr + ((i_b * T + t0 + o_c) * H + i_h)[:, None] * K + o_k[None, :]).to(tl.float32) b_g = tl.cumsum(tl.load(g_ptr + ((i_b * T + t0 + o_c) * H + i_h)[:, None] * K + o_k[None, :]), 0) b_A += tl.dot(b_k * tl.exp(b_g), tl.trans(b_k * tl.exp(-b_g))) b_k = tl.load(k_ptr + ((i_b * T + t0 + o_c) * H + i_h)[:, None] * K + (o_k + BK)[None, :]).to(tl.float32) b_g = tl.cumsum(tl.load(g_ptr + ((i_b * T + t0 + o_c) * H + i_h)[:, None] * K + (o_k + BK)[None, :]), 0) b_A += tl.dot(b_k * tl.exp(b_g), tl.trans(b_k * tl.exp(-b_g))) b_beta = tl.load(beta_ptr + (i_b * T + t0 + o_c) * H + i_h).to(tl.float32) b_A = tl.where(o_c[:, None] > o_c[None, :], b_A * b_beta[:, None], 0.0) nid = i_bh * NT + i_t tl.store(A_ptr + nid * BT * BT + o_c[:, None] * BT + o_c[None, :], b_A) @triton.jit def _invert_A_kernel( A_ptr, beta_ptr, T, H, NT, BT: tl.constexpr, ): """A <- (I+A)^{-1} * diag(beta). A strictly lower, packed [BH*NT, BT, BT].""" i_t = tl.program_id(0) i_bh = tl.program_id(1) i_b = i_bh // H i_h = i_bh % H nid = i_bh * NT + i_t t0 = i_t * BT A = A_ptr + nid * BT * BT o = tl.arange(0, 16) m_A = o[:, None] > o[None, :] m_I = o[:, None] == o[None, :] b11 = -tl.where(m_A, tl.load(A + o[:, None] * BT + o[None, :]).to(tl.float32), 0.0) b22 = -tl.where(m_A, tl.load(A + (16 + o)[:, None] * BT + (16 + o)[None, :]).to(tl.float32), 0.0) b33 = -tl.where(m_A, tl.load(A + (32 + o)[:, None] * BT + (32 + o)[None, :]).to(tl.float32), 0.0) b44 = -tl.where(m_A, tl.load(A + (48 + o)[:, None] * BT + (48 + o)[None, :]).to(tl.float32), 0.0) for i in range(2, 16): a = -tl.load(A + i * BT + o) a = tl.where(o < i, a, 0.0) a += tl.sum(a[:, None] * b11, 0) b11 = tl.where((o == i)[:, None], a[None, :], b11) for i in range(2, 16): a = -tl.load(A + (16 + i) * BT + (16 + o)) a = tl.where(o < i, a, 0.0) a += tl.sum(a[:, None] * b22, 0) b22 = tl.where((o == i)[:, None], a[None, :], b22) for i in range(2, 16): a = -tl.load(A + (32 + i) * BT + (32 + o)) a = tl.where(o < i, a, 0.0) a += tl.sum(a[:, None] * b33, 0) b33 = tl.where((o == i)[:, None], a[None, :], b33) for i in range(2, 16): a = -tl.load(A + (48 + i) * BT + (48 + o)) a = tl.where(o < i, a, 0.0) a += tl.sum(a[:, None] * b44, 0) b44 = tl.where((o == i)[:, None], a[None, :], b44) b11 += m_I b22 += m_I b33 += m_I b44 += m_I a21 = tl.load(A + (16 + o)[:, None] * BT + o[None, :]).to(tl.float32) a31 = tl.load(A + (32 + o)[:, None] * BT + o[None, :]).to(tl.float32) a32 = tl.load(A + (32 + o)[:, None] * BT + (16 + o)[None, :]).to(tl.float32) a41 = tl.load(A + (48 + o)[:, None] * BT + o[None, :]).to(tl.float32) a42 = tl.load(A + (48 + o)[:, None] * BT + (16 + o)[None, :]).to(tl.float32) a43 = tl.load(A + (48 + o)[:, None] * BT + (32 + o)[None, :]).to(tl.float32) ai21 = -tl.dot(tl.dot(b22, a21), b11) ai32 = -tl.dot(tl.dot(b33, a32), b22) ai43 = -tl.dot(tl.dot(b44, a43), b33) ai31 = -tl.dot(b33, tl.dot(a31, b11) + tl.dot(a32, ai21)) ai42 = -tl.dot(b44, tl.dot(a42, b22) + tl.dot(a43, ai32)) ai41 = -tl.dot(b44, tl.dot(a41, b11) + tl.dot(a42, ai21) + tl.dot(a43, ai31)) b0 = tl.load(beta_ptr + (i_b * T + t0 + o) * H + i_h).to(tl.float32) b1 = tl.load(beta_ptr + (i_b * T + t0 + 16 + o) * H + i_h).to(tl.float32) b2 = tl.load(beta_ptr + (i_b * T + t0 + 32 + o) * H + i_h).to(tl.float32) b3 = tl.load(beta_ptr + (i_b * T + t0 + 48 + o) * H + i_h).to(tl.float32) tl.store(A + o[:, None] * BT + o[None, :], b11 * b0[None, :]) tl.store(A + (16 + o)[:, None] * BT + (16 + o)[None, :], b22 * b1[None, :]) tl.store(A + (32 + o)[:, None] * BT + (32 + o)[None, :], b33 * b2[None, :]) tl.store(A + (48 + o)[:, None] * BT + (48 + o)[None, :], b44 * b3[None, :]) tl.store(A + (16 + o)[:, None] * BT + o[None, :], ai21 * b0[None, :]) tl.store(A + (32 + o)[:, None] * BT + (16 + o)[None, :], ai32 * b1[None, :]) tl.store(A + (48 + o)[:, None] * BT + (32 + o)[None, :], ai43 * b2[None, :]) tl.store(A + (32 + o)[:, None] * BT + o[None, :], ai31 * b0[None, :]) tl.store(A + (48 + o)[:, None] * BT + (16 + o)[None, :], ai42 * b1[None, :]) tl.store(A + (48 + o)[:, None] * BT + o[None, :], ai41 * b0[None, :]) @triton.jit def _build_wy_kernel( q_ptr, k_ptr, v_ptr, g_ptr, A_ptr, w_ptr, u_ptr, aqk_ptr, qg_ptr, kbar_ptr, glast_ptr, scale, T, H, NT, BT: tl.constexpr, BK: tl.constexpr, ): i_t = tl.program_id(0) i_bh = tl.program_id(1) i_b = i_bh // H i_h = i_bh % H t0 = i_t * BT K = 128 V = 128 o_c = tl.arange(0, BT) o_k = tl.arange(0, BK) nid = i_bh * NT + i_t pack = nid * BT b_A = tl.load(A_ptr + nid * BT * BT + o_c[:, None] * BT + o_c[None, :]) b_Aqk = tl.zeros([BT, BT], dtype=tl.float32) b_k = tl.load(k_ptr + ((i_b * T + t0 + o_c) * H + i_h)[:, None] * K + o_k[None, :]).to(tl.float32) b_g = tl.cumsum(tl.load(g_ptr + ((i_b * T + t0 + o_c) * H + i_h)[:, None] * K + o_k[None, :]), 0) b_eg = tl.exp(b_g) b_gl = tl.sum(tl.where(o_c[:, None] == (BT - 1), b_g, 0.0), axis=0) tl.store(glast_ptr + nid * K + o_k, b_gl) b_q = tl.load(q_ptr + ((i_b * T + t0 + o_c) * H + i_h)[:, None] * K + o_k[None, :]).to(tl.float32) b_qg = b_q * scale * b_eg tl.store(w_ptr + (pack + o_c)[:, None] * K + o_k[None, :], tl.dot(b_A, b_k * b_eg).to(tl.bfloat16)) tl.store(qg_ptr + (pack + o_c)[:, None] * K + o_k[None, :], b_qg.to(tl.bfloat16)) tl.store(kbar_ptr + (pack + o_c)[:, None] * K + o_k[None, :], (b_k * tl.exp(b_gl[None, :] - b_g)).to(tl.bfloat16)) b_Aqk += tl.dot(b_qg, tl.trans(b_k * tl.exp(-b_g))) b_k = tl.load(k_ptr + ((i_b * T + t0 + o_c) * H + i_h)[:, None] * K + (o_k + BK)[None, :]).to(tl.float32) b_g = tl.cumsum(tl.load(g_ptr + ((i_b * T + t0 + o_c) * H + i_h)[:, None] * K + (o_k + BK)[None, :]), 0) b_eg = tl.exp(b_g) b_gl = tl.sum(tl.where(o_c[:, None] == (BT - 1), b_g, 0.0), axis=0) tl.store(glast_ptr + nid * K + (o_k + BK), b_gl) b_q = tl.load(q_ptr + ((i_b * T + t0 + o_c) * H + i_h)[:, None] * K + (o_k + BK)[None, :]).to(tl.float32) b_qg = b_q * scale * b_eg tl.store(w_ptr + (pack + o_c)[:, None] * K + (o_k + BK)[None, :], tl.dot(b_A, b_k * b_eg).to(tl.bfloat16)) tl.store(qg_ptr + (pack + o_c)[:, None] * K + (o_k + BK)[None, :], b_qg.to(tl.bfloat16)) tl.store(kbar_ptr + (pack + o_c)[:, None] * K + (o_k + BK)[None, :], (b_k * tl.exp(b_gl[None, :] - b_g)).to(tl.bfloat16)) b_Aqk += tl.dot(b_qg, tl.trans(b_k * tl.exp(-b_g))) b_Aqk = tl.where(o_c[:, None] >= o_c[None, :], b_Aqk, 0.0) tl.store(aqk_ptr + (pack + o_c)[:, None] * BT + o_c[None, :], b_Aqk.to(tl.bfloat16)) b_v = tl.load(v_ptr + ((i_b * T + t0 + o_c) * H + i_h)[:, None] * V + o_k[None, :]).to(tl.float32) tl.store(u_ptr + (pack + o_c)[:, None] * V + o_k[None, :], tl.dot(b_A, b_v).to(tl.bfloat16)) b_v = tl.load(v_ptr + ((i_b * T + t0 + o_c) * H + i_h)[:, None] * V + (o_k + BK)[None, :]).to(tl.float32) tl.store(u_ptr + (pack + o_c)[:, None] * V + (o_k + BK)[None, :], tl.dot(b_A, b_v).to(tl.bfloat16)) @triton.jit(do_not_specialize=["NT"]) def _inter_kernel( w_ptr, u_ptr, aqk_ptr, qg_ptr, kbar_ptr, glast_ptr, o_ptr, T, H, NT, BT: tl.constexpr, BK: tl.constexpr, BV: tl.constexpr, ): i_bh = tl.program_id(0) iv = tl.program_id(1) i_b = i_bh // H i_h = i_bh % H o_c = tl.arange(0, BT) o_k = tl.arange(0, BK) o_v = iv * BV + tl.arange(0, BV) K = 128 V = 128 s0 = tl.zeros([BK, BV], dtype=tl.float32) s1 = tl.zeros([BK, BV], dtype=tl.float32) s2 = tl.zeros([BK, BV], dtype=tl.float32) s3 = tl.zeros([BK, BV], dtype=tl.float32) for it in range(0, NT): idx = (i_bh * NT + it) * BT t0 = it * BT b_vi = tl.load(u_ptr + (idx + o_c)[:, None] * V + o_v[None, :]).to(tl.float32) b_w = tl.load(w_ptr + (idx + o_c)[:, None] * K + o_k[None, :]) b_vi -= tl.dot(b_w, s0.to(tl.bfloat16)) b_w = tl.load(w_ptr + (idx + o_c)[:, None] * K + (o_k + BK)[None, :]) b_vi -= tl.dot(b_w, s1.to(tl.bfloat16)) b_w = tl.load(w_ptr + (idx + o_c)[:, None] * K + (o_k + 2 * BK)[None, :]) b_vi -= tl.dot(b_w, s2.to(tl.bfloat16)) b_w = tl.load(w_ptr + (idx + o_c)[:, None] * K + (o_k + 3 * BK)[None, :]) b_vi -= tl.dot(b_w, s3.to(tl.bfloat16)) b_o = tl.zeros([BT, BV], dtype=tl.float32) b_q = tl.load(qg_ptr + (idx + o_c)[:, None] * K + o_k[None, :]) b_o += tl.dot(b_q, s0.to(tl.bfloat16)) b_q = tl.load(qg_ptr + (idx + o_c)[:, None] * K + (o_k + BK)[None, :]) b_o += tl.dot(b_q, s1.to(tl.bfloat16)) b_q = tl.load(qg_ptr + (idx + o_c)[:, None] * K + (o_k + 2 * BK)[None, :]) b_o += tl.dot(b_q, s2.to(tl.bfloat16)) b_q = tl.load(qg_ptr + (idx + o_c)[:, None] * K + (o_k + 3 * BK)[None, :]) b_o += tl.dot(b_q, s3.to(tl.bfloat16)) b_A = tl.load(aqk_ptr + (idx + o_c)[:, None] * BT + o_c[None, :]) b_o += tl.dot(b_A, b_vi.to(tl.bfloat16)) tl.store(o_ptr + ((i_b * T + t0 + o_c) * H + i_h)[:, None] * V + o_v[None, :], b_o.to(tl.bfloat16)) b_g = tl.load(glast_ptr + (i_bh * NT + it) * K + o_k) s0 = s0 * tl.exp(b_g)[:, None] b_kb = tl.load(kbar_ptr + (idx + o_c)[:, None] * K + o_k[None, :]) s0 += tl.dot(tl.trans(b_kb), b_vi.to(tl.bfloat16)) b_g = tl.load(glast_ptr + (i_bh * NT + it) * K + (o_k + BK)) s1 = s1 * tl.exp(b_g)[:, None] b_kb = tl.load(kbar_ptr + (idx + o_c)[:, None] * K + (o_k + BK)[None, :]) s1 += tl.dot(tl.trans(b_kb), b_vi.to(tl.bfloat16)) b_g = tl.load(glast_ptr + (i_bh * NT + it) * K + (o_k + 2 * BK)) s2 = s2 * tl.exp(b_g)[:, None] b_kb = tl.load(kbar_ptr + (idx + o_c)[:, None] * K + (o_k + 2 * BK)[None, :]) s2 += tl.dot(tl.trans(b_kb), b_vi.to(tl.bfloat16)) b_g = tl.load(glast_ptr + (i_bh * NT + it) * K + (o_k + 3 * BK)) s3 = s3 * tl.exp(b_g)[:, None] b_kb = tl.load(kbar_ptr + (idx + o_c)[:, None] * K + (o_k + 3 * BK)[None, :]) s3 += tl.dot(tl.trans(b_kb), b_vi.to(tl.bfloat16)) _WS: dict = {} def _workspace(B, T, H, Kdim, Vdim, BT, device): key = (B, T, H, Kdim, Vdim, BT, device) ws = _WS.get(key) if ws is not None: return ws NT = T // BT N = B * H * NT ws = { "A": torch.empty(N, BT, BT, device=device, dtype=torch.float32), "w": torch.empty(N * BT, Kdim, device=device, dtype=torch.bfloat16), "u": torch.empty(N * BT, Vdim, device=device, dtype=torch.bfloat16), "aqk": torch.empty(N * BT, BT, device=device, dtype=torch.bfloat16), "qg": torch.empty(N * BT, Kdim, device=device, dtype=torch.bfloat16), "kbar": torch.empty(N * BT, Kdim, device=device, dtype=torch.bfloat16), "glast": torch.empty(N, Kdim, device=device, dtype=torch.float32), "o": torch.empty(B, T, H, Vdim, device=device, dtype=torch.bfloat16), } _WS[key] = ws return ws def _kda_forward(q, k, v, g, beta, scale, chunk_size): B, T, H, Kdim = q.shape Vdim = v.shape[-1] BT = chunk_size NT = T // BT q = q.contiguous() k = k.contiguous() v = v.contiguous() g = g.contiguous() beta = beta.contiguous() ws = _workspace(B, T, H, Kdim, Vdim, BT, q.device) grid = (NT, B * H) _build_A_kernel[grid](k, g, beta, ws["A"], T, H, NT, BT=BT, BK=64, num_warps=4, num_stages=2) _invert_A_kernel[grid](ws["A"], beta, T, H, NT, BT=BT, num_warps=4, num_stages=1) _build_wy_kernel[grid]( q, k, v, g, ws["A"], ws["w"], ws["u"], ws["aqk"], ws["qg"], ws["kbar"], ws["glast"], float(scale), T, H, NT, BT=BT, BK=64, num_warps=4, num_stages=2, ) _inter_kernel[(B * H, 4)]( ws["w"], ws["u"], ws["aqk"], ws["qg"], ws["kbar"], ws["glast"], ws["o"], T, H, NT, BT=BT, BK=32, BV=32, num_warps=4, num_stages=2, ) return ws["o"] 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) self._graph = None self._graph_ptrs = None def forward(self, q, k, v, g, beta): q = q.contiguous() k = k.contiguous() v = v.contiguous() g = g.contiguous() beta = beta.contiguous() ptrs = (q.data_ptr(), k.data_ptr(), v.data_ptr(), g.data_ptr(), beta.data_ptr()) if self._graph is None or ptrs != self._graph_ptrs: out = _kda_forward(q, k, v, g, beta, scale=self.scale, chunk_size=self.chunk_size) try: _kda_forward(q, k, v, g, beta, scale=self.scale, chunk_size=self.chunk_size) torch.cuda.synchronize() graph = torch.cuda.CUDAGraph() with torch.cuda.graph(graph): _kda_forward(q, k, v, g, beta, scale=self.scale, chunk_size=self.chunk_size) self._graph = graph self._graph_ptrs = ptrs except Exception: self._graph = None self._graph_ptrs = None return out self._graph.replay() ws = _workspace(q.shape[0], q.shape[1], q.shape[2], q.shape[3], v.shape[-1], self.chunk_size, q.device) return ws["o"] 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]