"""SM90-oriented Triton forward kernels for Kimi Delta Attention. Each 64-token chunk is prepared independently with tensor-core products and a four-block triangular inverse. A final persistent kernel owns one (batch, head, value-tile), fuses the inter-chunk recurrence with the output projection, and keeps its 128 x BV state on chip for the complete sequence. """ from __future__ import annotations import torch import torch.nn as nn import triton import triton.language as tl PAIR_WARPS = 4 PAIR_STAGES = 3 SOLVE_WARPS = 2 SOLVE_STAGES = 1 PROJECT_WARPS = 4 PROJECT_STAGES = 3 STATE_BV = 16 STATE_WARPS = 4 OUTPUT_BV = 128 OUTPUT_WARPS = 4 OUTPUT_STAGES = 3 @triton.jit def _delta_forward_kernel( q_ptr, k_ptr, v_ptr, g_ptr, beta_ptr, out_ptr, T, H: tl.constexpr, D: tl.constexpr, BV: tl.constexpr, SCALE: tl.constexpr, ): pid_v = tl.program_id(0) pid_bh = tl.program_id(1) b = pid_bh // H h = pid_bh - b * H ik = tl.arange(0, D) iv = pid_v * BV + tl.arange(0, BV) mv = iv < D # This tensor is distributed over the participating warps and stays in # registers for the dynamic sequence loop. state = tl.zeros((D, BV), tl.float32) t = 0 while t < T: qk_base = ((b * T + t) * H + h) * D v_base = qk_base bh_base = (b * T + t) * H + h q = tl.load(q_ptr + qk_base + ik).to(tl.float32) k = tl.load(k_ptr + qk_base + ik).to(tl.float32) decay = tl.exp(tl.load(g_ptr + qk_base + ik)) value = tl.load(v_ptr + v_base + iv, mask=mv, other=0.0).to(tl.float32) beta = tl.load(beta_ptr + bh_base).to(tl.float32) state *= decay[:, None] prediction = tl.sum(state * k[:, None], axis=0) residual = beta * (value - prediction) state += k[:, None] * residual[None, :] output = tl.sum(state * q[:, None], axis=0) * SCALE tl.store(out_ptr + v_base + iv, output, mask=mv) t += 1 @triton.jit def _local_prefix_kernel( src, dst, T, H: tl.constexpr, D: tl.constexpr, BT: tl.constexpr, BS: tl.constexpr, ): pid_s = tl.program_id(0) pid_t = tl.program_id(1) pid_bh = tl.program_id(2) b = pid_bh // H h = pid_bh - b * H base = (b * T * H + h) * D p_src = tl.make_block_ptr( src + base, (T, D), (H * D, 1), (pid_t * BT, pid_s * BS), (BT, BS), (1, 0), ) p_dst = tl.make_block_ptr( dst + base, (T, D), (H * D, 1), (pid_t * BT, pid_s * BS), (BT, BS), (1, 0), ) x = tl.load(p_src) tl.store(p_dst, tl.cumsum(x, axis=0)) @triton.jit def _build_chunks_kernel( q_ptr, k_ptr, v_ptr, gc_ptr, beta_ptr, score_ptr, w_ptr, u_ptr, kg_ptr, T, H: tl.constexpr, D: tl.constexpr, BT: tl.constexpr, SCALE: tl.constexpr, ): pid_t = tl.program_id(0) pid_bh = tl.program_id(1) b = pid_bh // H h = pid_bh - b * H tc = pid_t * BT qk_base = (b * T * H + h) * D s_base = (b * T * H + h) * BT beta_base = b * T * H + h p_q = tl.make_block_ptr(q_ptr + qk_base, (T, D), (H * D, 1), (tc, 0), (BT, D), (1, 0)) p_k = tl.make_block_ptr(k_ptr + qk_base, (T, D), (H * D, 1), (tc, 0), (BT, D), (1, 0)) p_v = tl.make_block_ptr(v_ptr + qk_base, (T, D), (H * D, 1), (tc, 0), (BT, D), (1, 0)) p_g = tl.make_block_ptr(gc_ptr + qk_base, (T, D), (H * D, 1), (tc, 0), (BT, D), (1, 0)) p_beta = tl.make_block_ptr(beta_ptr + beta_base, (T,), (H,), (tc,), (BT,), (0,)) q = tl.load(p_q) k = tl.load(p_k) g = tl.load(p_g) beta = tl.load(p_beta).to(tl.float32) # Centering the exponentials does not change either product and keeps both # factors well-scaled even for an unusually lopsided chunk. anchor = tl.sum(tl.where(tl.arange(0, BT)[:, None] == BT // 2, g, 0.0), axis=0) ep = tl.exp(g - anchor[None, :]) em = tl.exp(anchor[None, :] - g) qg = q.to(tl.float32) * ep kg_left = k.to(tl.float32) * ep kg_right = k.to(tl.float32) * em scores = tl.dot(qg, tl.trans(kg_right), input_precision="tf32") * SCALE interaction = tl.dot(kg_left, tl.trans(kg_right), input_precision="tf32") ii = tl.arange(0, BT) lower = ii[:, None] > ii[None, :] causal = ii[:, None] >= ii[None, :] interaction = tl.where(lower, interaction * beta[:, None], 0.0) scores = tl.where(causal, scores, 0.0) # Invert the unit lower-triangular system by forward substitution. Each # completed row is immediately visible to subsequent rows in `inv`. inv = -interaction for row_idx in range(2, BT): row = -tl.sum( tl.where(ii[:, None] == row_idx, interaction, 0.0), axis=0 ) row += tl.sum(row[:, None] * inv, axis=0) inv = tl.where((ii == row_idx)[:, None], row[None, :], inv) inv += (ii[:, None] == ii[None, :]) transform = inv * beta[None, :] p_score = tl.make_block_ptr( score_ptr + s_base, (T, BT), (H * BT, 1), (tc, 0), (BT, BT), (1, 0) ) tl.store(p_score, scores.to(tl.bfloat16)) # These three intermediates are deliberately bf16: their consumers are # tensor-core products and the fp32 state holds the inter-chunk accuracy. v = tl.load(p_v) u = tl.dot(transform.to(tl.bfloat16), v) p_u = tl.make_block_ptr(u_ptr + qk_base, (T, D), (H * D, 1), (tc, 0), (BT, D), (1, 0)) tl.store(p_u, u.to(tl.bfloat16)) k_exp = (k.to(tl.float32) * tl.exp(g)).to(tl.bfloat16) w = tl.dot(transform.to(tl.bfloat16), k_exp) p_w = tl.make_block_ptr(w_ptr + qk_base, (T, D), (H * D, 1), (tc, 0), (BT, D), (1, 0)) tl.store(p_w, w.to(tl.bfloat16)) g_last = tl.sum(tl.where(ii[:, None] == BT - 1, g, 0.0), axis=0) k_tail = k.to(tl.float32) * tl.exp(g_last[None, :] - g) p_kg = tl.make_block_ptr(kg_ptr + qk_base, (T, D), (H * D, 1), (tc, 0), (BT, D), (1, 0)) tl.store(p_kg, k_tail.to(tl.bfloat16)) @triton.jit def _pair_matrices_kernel( q_ptr, k_ptr, g_ptr, gc_ptr, beta_ptr, score_ptr, lower_ptr, T, H: tl.constexpr, D: tl.constexpr, BT: tl.constexpr, SCALE: tl.constexpr, ): pid_t = tl.program_id(0) pid_bh = tl.program_id(1) b = pid_bh // H h = pid_bh - b * H tc = pid_t * BT qk_base = (b * T * H + h) * D m_base = (b * T * H + h) * BT beta_base = b * T * H + h p_q = tl.make_block_ptr(q_ptr + qk_base, (T, D), (H * D, 1), (tc, 0), (BT, D), (1, 0)) p_k = tl.make_block_ptr(k_ptr + qk_base, (T, D), (H * D, 1), (tc, 0), (BT, D), (1, 0)) p_g = tl.make_block_ptr(g_ptr + qk_base, (T, D), (H * D, 1), (tc, 0), (BT, D), (1, 0)) p_gc = tl.make_block_ptr(gc_ptr + qk_base, (T, D), (H * D, 1), (tc, 0), (BT, D), (1, 0)) p_b = tl.make_block_ptr(beta_ptr + beta_base, (T,), (H,), (tc,), (BT,), (0,)) q = tl.load(p_q) k = tl.load(p_k) g = tl.cumsum(tl.load(p_g), axis=0) tl.store(p_gc, g) beta = tl.load(p_b).to(tl.float32) rows = tl.arange(0, BT) anchor = tl.sum(tl.where(rows[:, None] == BT // 2, g, 0.0), axis=0) ep = tl.exp(g - anchor[None, :]) em = tl.exp(anchor[None, :] - g) right = k.to(tl.float32) * em score = tl.dot(q.to(tl.float32) * ep, tl.trans(right), input_precision="tf32") * SCALE lower = tl.dot(k.to(tl.float32) * ep, tl.trans(right), input_precision="tf32") score = tl.where(rows[:, None] >= rows[None, :], score, 0.0) lower = tl.where(rows[:, None] > rows[None, :], lower * beta[:, None], 0.0) p_score = tl.make_block_ptr(score_ptr + m_base, (T, BT), (H * BT, 1), (tc, 0), (BT, BT), (1, 0)) p_lower = tl.make_block_ptr(lower_ptr + m_base, (T, BT), (H * BT, 1), (tc, 0), (BT, BT), (1, 0)) tl.store(p_score, score.to(tl.bfloat16)) tl.store(p_lower, lower) @triton.jit def _load_matrix16(ptr, base, T, H: tl.constexpr, BT: tl.constexpr, BC: tl.constexpr, tc, RO: tl.constexpr, CO: tl.constexpr): p = tl.make_block_ptr( ptr + base, (T, BT), (H * BT, 1), (tc + RO * BC, CO * BC), (BC, BC), (1, 0), ) return tl.load(p) @triton.jit def _load_beta16(ptr, base, T, H: tl.constexpr, BC: tl.constexpr, tc, BLOCK: tl.constexpr): p = tl.make_block_ptr(ptr + base, (T,), (H,), (tc + BLOCK * BC,), (BC,), (0,)) return tl.load(p).to(tl.float32) @triton.jit def _store_matrix16(ptr, x, base, T, H: tl.constexpr, BT: tl.constexpr, BC: tl.constexpr, tc, RO: tl.constexpr, CO: tl.constexpr): p = tl.make_block_ptr( ptr + base, (T, BT), (H * BT, 1), (tc + RO * BC, CO * BC), (BC, BC), (1, 0), ) tl.store(p, x.to(tl.bfloat16)) @triton.jit def _solve64_kernel( lower_ptr, beta_ptr, transform_ptr, T, H: tl.constexpr, BT: tl.constexpr, BC: tl.constexpr, ): pid_t = tl.program_id(0) pid_bh = tl.program_id(1) b = pid_bh // H h = pid_bh - b * H tc = pid_t * BT m_base = (b * T * H + h) * BT beta_base = b * T * H + h l00 = _load_matrix16(lower_ptr, m_base, T, H, BT, BC, tc, 0, 0) l10 = _load_matrix16(lower_ptr, m_base, T, H, BT, BC, tc, 1, 0) l11 = _load_matrix16(lower_ptr, m_base, T, H, BT, BC, tc, 1, 1) l20 = _load_matrix16(lower_ptr, m_base, T, H, BT, BC, tc, 2, 0) l21 = _load_matrix16(lower_ptr, m_base, T, H, BT, BC, tc, 2, 1) l22 = _load_matrix16(lower_ptr, m_base, T, H, BT, BC, tc, 2, 2) l30 = _load_matrix16(lower_ptr, m_base, T, H, BT, BC, tc, 3, 0) l31 = _load_matrix16(lower_ptr, m_base, T, H, BT, BC, tc, 3, 1) l32 = _load_matrix16(lower_ptr, m_base, T, H, BT, BC, tc, 3, 2) l33 = _load_matrix16(lower_ptr, m_base, T, H, BT, BC, tc, 3, 3) ii = tl.arange(0, BC) strict = ii[:, None] > ii[None, :] ident = ii[:, None] == ii[None, :] a00 = -tl.where(strict, l00, 0.0) a11 = -tl.where(strict, l11, 0.0) a22 = -tl.where(strict, l22, 0.0) a33 = -tl.where(strict, l33, 0.0) for row_idx in range(2, BC): r0 = -tl.sum(tl.where(ii[:, None] == row_idx, l00, 0.0), axis=0) r1 = -tl.sum(tl.where(ii[:, None] == row_idx, l11, 0.0), axis=0) r2 = -tl.sum(tl.where(ii[:, None] == row_idx, l22, 0.0), axis=0) r3 = -tl.sum(tl.where(ii[:, None] == row_idx, l33, 0.0), axis=0) r0 += tl.sum(r0[:, None] * a00, axis=0) r1 += tl.sum(r1[:, None] * a11, axis=0) r2 += tl.sum(r2[:, None] * a22, axis=0) r3 += tl.sum(r3[:, None] * a33, axis=0) select = (ii == row_idx)[:, None] a00 = tl.where(select, r0[None, :], a00) a11 = tl.where(select, r1[None, :], a11) a22 = tl.where(select, r2[None, :], a22) a33 = tl.where(select, r3[None, :], a33) a00 += ident a11 += ident a22 += ident a33 += ident a10 = -tl.dot(tl.dot(a11, l10, input_precision="tf32"), a00, input_precision="tf32") a21 = -tl.dot(tl.dot(a22, l21, input_precision="tf32"), a11, input_precision="tf32") a20 = -tl.dot( a22, tl.dot(l20, a00, input_precision="tf32") + tl.dot(l21, a10, input_precision="tf32"), input_precision="tf32", ) a32 = -tl.dot(tl.dot(a33, l32, input_precision="tf32"), a22, input_precision="tf32") a31 = -tl.dot( a33, tl.dot(l31, a11, input_precision="tf32") + tl.dot(l32, a21, input_precision="tf32"), input_precision="tf32", ) a30 = -tl.dot( a33, tl.dot(l30, a00, input_precision="tf32") + tl.dot(l31, a10, input_precision="tf32") + tl.dot(l32, a20, input_precision="tf32"), input_precision="tf32", ) b0 = _load_beta16(beta_ptr, beta_base, T, H, BC, tc, 0) b1 = _load_beta16(beta_ptr, beta_base, T, H, BC, tc, 1) b2 = _load_beta16(beta_ptr, beta_base, T, H, BC, tc, 2) b3 = _load_beta16(beta_ptr, beta_base, T, H, BC, tc, 3) _store_matrix16(transform_ptr, a00 * b0[None, :], m_base, T, H, BT, BC, tc, 0, 0) _store_matrix16(transform_ptr, a10 * b0[None, :], m_base, T, H, BT, BC, tc, 1, 0) _store_matrix16(transform_ptr, a11 * b1[None, :], m_base, T, H, BT, BC, tc, 1, 1) _store_matrix16(transform_ptr, a20 * b0[None, :], m_base, T, H, BT, BC, tc, 2, 0) _store_matrix16(transform_ptr, a21 * b1[None, :], m_base, T, H, BT, BC, tc, 2, 1) _store_matrix16(transform_ptr, a22 * b2[None, :], m_base, T, H, BT, BC, tc, 2, 2) _store_matrix16(transform_ptr, a30 * b0[None, :], m_base, T, H, BT, BC, tc, 3, 0) _store_matrix16(transform_ptr, a31 * b1[None, :], m_base, T, H, BT, BC, tc, 3, 1) _store_matrix16(transform_ptr, a32 * b2[None, :], m_base, T, H, BT, BC, tc, 3, 2) _store_matrix16(transform_ptr, a33 * b3[None, :], m_base, T, H, BT, BC, tc, 3, 3) @triton.jit def _project_chunks_kernel( k_ptr, v_ptr, gc_ptr, transform_ptr, w_ptr, u_ptr, kg_ptr, T, H: tl.constexpr, D: tl.constexpr, BT: tl.constexpr, ): pid_t = tl.program_id(0) pid_bh = tl.program_id(1) b = pid_bh // H h = pid_bh - b * H tc = pid_t * BT qk_base = (b * T * H + h) * D m_base = (b * T * H + h) * BT p_a = tl.make_block_ptr(transform_ptr + m_base, (T, BT), (H * BT, 1), (tc, 0), (BT, BT), (1, 0)) a = tl.load(p_a) ii = tl.arange(0, BT) a = tl.where(ii[:, None] >= ii[None, :], a, 0.0) for block in range(2): d0 = block * 64 p_k = tl.make_block_ptr(k_ptr + qk_base, (T, D), (H * D, 1), (tc, d0), (BT, 64), (1, 0)) p_v = tl.make_block_ptr(v_ptr + qk_base, (T, D), (H * D, 1), (tc, d0), (BT, 64), (1, 0)) p_g = tl.make_block_ptr(gc_ptr + qk_base, (T, D), (H * D, 1), (tc, d0), (BT, 64), (1, 0)) k = tl.load(p_k) v = tl.load(p_v) g = tl.load(p_g) u = tl.dot(a, v) w = tl.dot(a, (k.to(tl.float32) * tl.exp(g)).to(tl.bfloat16)) last = tl.sum(tl.where(ii[:, None] == BT - 1, g, 0.0), axis=0) kg = k.to(tl.float32) * tl.exp(last[None, :] - g) p_u = tl.make_block_ptr(u_ptr + qk_base, (T, D), (H * D, 1), (tc, d0), (BT, 64), (1, 0)) p_w = tl.make_block_ptr(w_ptr + qk_base, (T, D), (H * D, 1), (tc, d0), (BT, 64), (1, 0)) p_kg = tl.make_block_ptr(kg_ptr + qk_base, (T, D), (H * D, 1), (tc, d0), (BT, 64), (1, 0)) tl.store(p_u, u.to(tl.bfloat16)) tl.store(p_w, w.to(tl.bfloat16)) tl.store(p_kg, kg.to(tl.bfloat16)) @triton.jit def _chunks_recurrence_kernel( q_ptr, gc_ptr, score_ptr, w_ptr, u_ptr, kg_ptr, out_ptr, T, H: tl.constexpr, D: tl.constexpr, BT: tl.constexpr, BV: tl.constexpr, SCALE: tl.constexpr, ): pid_v = tl.program_id(0) pid_bh = tl.program_id(1) b = pid_bh // H h = pid_bh - b * H v0 = pid_v * BV qk_base = (b * T * H + h) * D s_base = (b * T * H + h) * BT state0 = tl.zeros((64, BV), tl.float32) state1 = tl.zeros((64, BV), tl.float32) chunk = 0 while chunk < T // BT: tc = chunk * BT p_w0 = tl.make_block_ptr(w_ptr + qk_base, (T, D), (H * D, 1), (tc, 0), (BT, 64), (1, 0)) p_w1 = tl.make_block_ptr(w_ptr + qk_base, (T, D), (H * D, 1), (tc, 64), (BT, 64), (1, 0)) w0 = tl.load(p_w0) w1 = tl.load(p_w1) correction = tl.dot(w0, state0.to(tl.bfloat16)) correction += tl.dot(w1, state1.to(tl.bfloat16)) p_u = tl.make_block_ptr(u_ptr + qk_base, (T, D), (H * D, 1), (tc, v0), (BT, BV), (1, 0)) value = tl.load(p_u).to(tl.float32) - correction value_bf = value.to(tl.bfloat16) p_q0 = tl.make_block_ptr(q_ptr + qk_base, (T, D), (H * D, 1), (tc, 0), (BT, 64), (1, 0)) p_q1 = tl.make_block_ptr(q_ptr + qk_base, (T, D), (H * D, 1), (tc, 64), (BT, 64), (1, 0)) p_g0 = tl.make_block_ptr(gc_ptr + qk_base, (T, D), (H * D, 1), (tc, 0), (BT, 64), (1, 0)) p_g1 = tl.make_block_ptr(gc_ptr + qk_base, (T, D), (H * D, 1), (tc, 64), (BT, 64), (1, 0)) g0 = tl.load(p_g0) g1 = tl.load(p_g1) q0 = (tl.load(p_q0).to(tl.float32) * tl.exp(g0)).to(tl.bfloat16) q1 = (tl.load(p_q1).to(tl.float32) * tl.exp(g1)).to(tl.bfloat16) output = tl.dot(q0, state0.to(tl.bfloat16)) output += tl.dot(q1, state1.to(tl.bfloat16)) output *= SCALE p_score = tl.make_block_ptr( score_ptr + s_base, (T, BT), (H * BT, 1), (tc, 0), (BT, BT), (1, 0) ) score = tl.load(p_score) output += tl.dot(score, value_bf) p_out = tl.make_block_ptr(out_ptr + qk_base, (T, D), (H * D, 1), (tc, v0), (BT, BV), (1, 0)) tl.store(p_out, output.to(tl.bfloat16)) ik = tl.arange(0, 64) last0 = tl.load(gc_ptr + qk_base + (tc + BT - 1) * H * D + ik) last1 = tl.load(gc_ptr + qk_base + (tc + BT - 1) * H * D + 64 + ik) state0 *= tl.exp(last0)[:, None] state1 *= tl.exp(last1)[:, None] p_kg0 = tl.make_block_ptr(kg_ptr + qk_base, (D, T), (1, H * D), (0, tc), (64, BT), (0, 1)) p_kg1 = tl.make_block_ptr(kg_ptr + qk_base, (D, T), (1, H * D), (64, tc), (64, BT), (0, 1)) state0 += tl.dot(tl.load(p_kg0), value_bf) state1 += tl.dot(tl.load(p_kg1), value_bf) chunk += 1 @triton.jit def _chunks_recurrence128_kernel( q_ptr, gc_ptr, score_ptr, w_ptr, u_ptr, kg_ptr, out_ptr, T, H: tl.constexpr, D: tl.constexpr, BT: tl.constexpr, BV: tl.constexpr, SCALE: tl.constexpr, ): pid_v = tl.program_id(0) pid_bh = tl.program_id(1) b = pid_bh // H h = pid_bh - b * H v0 = pid_v * BV qk_base = (b * T * H + h) * D s_base = (b * T * H + h) * BT state = tl.zeros((D, BV), tl.float32) chunk = 0 while chunk < T // BT: tc = chunk * BT p_w = tl.make_block_ptr(w_ptr + qk_base, (T, D), (H * D, 1), (tc, 0), (BT, D), (1, 0)) p_u = tl.make_block_ptr(u_ptr + qk_base, (T, D), (H * D, 1), (tc, v0), (BT, BV), (1, 0)) value = tl.load(p_u).to(tl.float32) - tl.dot(tl.load(p_w), state.to(tl.bfloat16)) value_bf = value.to(tl.bfloat16) p_q = tl.make_block_ptr(q_ptr + qk_base, (T, D), (H * D, 1), (tc, 0), (BT, D), (1, 0)) p_g = tl.make_block_ptr(gc_ptr + qk_base, (T, D), (H * D, 1), (tc, 0), (BT, D), (1, 0)) g = tl.load(p_g) qg = (tl.load(p_q).to(tl.float32) * tl.exp(g)).to(tl.bfloat16) output = tl.dot(qg, state.to(tl.bfloat16)) * SCALE p_score = tl.make_block_ptr(score_ptr + s_base, (T, BT), (H * BT, 1), (tc, 0), (BT, BT), (1, 0)) output += tl.dot(tl.load(p_score), value_bf) p_out = tl.make_block_ptr(out_ptr + qk_base, (T, D), (H * D, 1), (tc, v0), (BT, BV), (1, 0)) tl.store(p_out, output.to(tl.bfloat16)) ik = tl.arange(0, D) last = tl.load(gc_ptr + qk_base + (tc + BT - 1) * H * D + ik) state *= tl.exp(last)[:, None] p_kg = tl.make_block_ptr(kg_ptr + qk_base, (D, T), (1, H * D), (0, tc), (D, BT), (0, 1)) state += tl.dot(tl.load(p_kg), value_bf) chunk += 1 @triton.jit def _state_chunks_kernel( gc_ptr, w_ptr, u_ptr, kg_ptr, states_ptr, values_ptr, T, H: tl.constexpr, D: tl.constexpr, BT: 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 v0 = pid_v * BV nt = T // BT qk_base = (b * T * H + h) * D state_base = (b * nt * H + h) * D * D state = tl.zeros((D, BV), tl.float32) chunk = 0 while chunk < nt: tc = chunk * BT p_state = tl.make_block_ptr( states_ptr + state_base + chunk * H * D * D, (D, D), (D, 1), (0, v0), (D, BV), (1, 0), ) tl.store(p_state, state.to(tl.bfloat16)) p_w = tl.make_block_ptr(w_ptr + qk_base, (T, D), (H * D, 1), (tc, 0), (BT, D), (1, 0)) p_u = tl.make_block_ptr(u_ptr + qk_base, (T, D), (H * D, 1), (tc, v0), (BT, BV), (1, 0)) value = tl.load(p_u).to(tl.float32) - tl.dot(tl.load(p_w), state.to(tl.bfloat16)) value_bf = value.to(tl.bfloat16) p_value = tl.make_block_ptr(values_ptr + qk_base, (T, D), (H * D, 1), (tc, v0), (BT, BV), (1, 0)) tl.store(p_value, value_bf) ik = tl.arange(0, D) last = tl.load(gc_ptr + qk_base + (tc + BT - 1) * H * D + ik) state *= tl.exp(last)[:, None] p_kg = tl.make_block_ptr(kg_ptr + qk_base, (D, T), (1, H * D), (0, tc), (D, BT), (0, 1)) state += tl.dot(tl.load(p_kg), value_bf) chunk += 1 @triton.jit def _output_chunks_kernel( q_ptr, gc_ptr, score_ptr, states_ptr, values_ptr, out_ptr, T, H: tl.constexpr, D: tl.constexpr, BT: tl.constexpr, BV: tl.constexpr, SCALE: tl.constexpr, ): pid_v = tl.program_id(0) chunk = tl.program_id(1) pid_bh = tl.program_id(2) b = pid_bh // H h = pid_bh - b * H v0 = pid_v * BV nt = T // BT tc = chunk * BT qk_base = (b * T * H + h) * D score_base = (b * T * H + h) * BT state_base = ((b * nt + chunk) * H + h) * D * D p_q = tl.make_block_ptr(q_ptr + qk_base, (T, D), (H * D, 1), (tc, 0), (BT, D), (1, 0)) p_g = tl.make_block_ptr(gc_ptr + qk_base, (T, D), (H * D, 1), (tc, 0), (BT, D), (1, 0)) p_state = tl.make_block_ptr(states_ptr + state_base, (D, D), (D, 1), (0, v0), (D, BV), (1, 0)) qg = (tl.load(p_q).to(tl.float32) * tl.exp(tl.load(p_g))).to(tl.bfloat16) output = tl.dot(qg, tl.load(p_state)) * SCALE p_score = tl.make_block_ptr(score_ptr + score_base, (T, BT), (H * BT, 1), (tc, 0), (BT, BT), (1, 0)) p_value = tl.make_block_ptr(values_ptr + qk_base, (T, D), (H * D, 1), (tc, v0), (BT, BV), (1, 0)) output += tl.dot(tl.load(p_score), tl.load(p_value)) p_out = tl.make_block_ptr(out_ptr + qk_base, (T, D), (H * D, 1), (tc, v0), (BT, BV), (1, 0)) tl.store(p_out, output.to(tl.bfloat16)) def _forward(q, k, v, g, beta, scale: float): B, T, H, D = q.shape bt = 64 nt = T // bt gc = torch.empty_like(g) score = torch.empty((B, T, H, bt), device=q.device, dtype=q.dtype) w = torch.empty_like(k) u = torch.empty_like(v) kg = torch.empty_like(k) out = torch.empty_like(v) lower = torch.empty((B, T, H, bt), device=q.device, dtype=torch.float32) transform = torch.empty_like(score) _pair_matrices_kernel[(nt, B * H)]( q, k, g, gc, beta, score, lower, T, H=H, D=D, BT=bt, SCALE=scale, num_warps=PAIR_WARPS, num_stages=PAIR_STAGES, ) _solve64_kernel[(nt, B * H)]( lower, beta, transform, T, H=H, BT=bt, BC=16, num_warps=SOLVE_WARPS, num_stages=SOLVE_STAGES, ) _project_chunks_kernel[(nt, B * H)]( k, v, gc, transform, w, u, kg, T, H=H, D=D, BT=bt, num_warps=PROJECT_WARPS, num_stages=PROJECT_STAGES, ) states = torch.empty((B, nt, H, D, D), device=q.device, dtype=q.dtype) values = torch.empty_like(v) long_context = T >= 4096 state_bv = 32 if long_context else STATE_BV state_warps = 8 if long_context else STATE_WARPS state_stages = 4 if long_context else 5 _state_chunks_kernel[(D // state_bv, B * H)]( gc, w, u, kg, states, values, T, H=H, D=D, BT=bt, BV=state_bv, num_warps=state_warps, num_stages=state_stages, ) output_bv = OUTPUT_BV _output_chunks_kernel[(D // output_bv, nt, B * H)]( q, gc, score, states, values, out, T, H=H, D=D, BT=bt, BV=output_bv, SCALE=scale, num_warps=OUTPUT_WARPS, num_stages=OUTPUT_STAGES, ) return out class Model(nn.Module): def __init__(self, B: int, T: int, H: int, K: int, V: int, chunk_size: int = 64): super().__init__() if K != 128 or V != 128 or chunk_size != 64: raise ValueError("This kernel is specialized for K=V=128 and chunks of 64") self.scale = float(K) ** -0.5 self.register_buffer("_dummy", torch.zeros(1), persistent=False) def forward(self, q, k, v, g, beta): return _forward(q, k, v, g, beta, self.scale) 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]