"""Kimi Delta Attention (KDA) forward, chunk form — custom Triton kernels for SM120. Two-kernel decomposition. Per chunk (BT=64, K=V=128; T % 64 == 0 so no masks): K1 (grid NT x B*H, chunk-parallel) g -> in-chunk fp32 cumsum A0 = (k e^{g-g_l}) (k e^{g_l-g})^T (bf16 tensor-core dot) B = tril(beta_row * A0, -1); M = (I+B)^-1 via exact nilpotent doubling u = M (beta v), w = M (beta e^g k) P' = diag(e^{g_l}) - (k e^{g_l-g})^T w [the chunk's state operator] C = (k e^{g_l-g})^T u Aqk = scale * tril(q e^{g-g_l} (k e^{g_l-g})^T, incl diag) qe = scale q e^g - Aqk w, t = Aqk u stores: P', C (B,H,NT,128,128 bf16), qe, t (token-major bf16) K2 (grid V/BV x B*H, sequential over NT, fused output) S = C_i + P'_i @ S (one dot per step; decay folded into P') o_i = qe_i @ S_i + t_i (uses the pre-update state) o = qg @ S_i + Aqk (u_i - w_i S_i) = (qg - Aqk w) S_i + Aqk u (exact identity). """ from __future__ import annotations import torch import torch.nn as nn import os import triton import triton.language as tl RCP_LN2 = tl.constexpr(1.4426950408889634) @triton.jit def _col2(x, AX: tl.constexpr): """[AX, 128] -> two [AX, 64] column halves.""" return tl.split(tl.permute(tl.reshape(x, (AX, 2, 64)), (0, 2, 1))) @triton.jit def _row2(x, AY: tl.constexpr): """[64, AY] -> two [32, AY] row halves.""" return tl.split(tl.permute(tl.reshape(x, (2, 32, AY)), (1, 2, 0))) @triton.jit def _col2b(x, AX: tl.constexpr): """[AX, 64] -> two [AX, 32] column halves.""" return tl.split(tl.permute(tl.reshape(x, (AX, 2, 32)), (0, 2, 1))) # --------------------------------------------------------------------------- # K1: chunk-local preparation (full-tile form, no layout shuffles) # --------------------------------------------------------------------------- @triton.jit def _kda_k1( q, k, v, g, beta, P, C, Pd, qe, t, T, H: tl.constexpr, NT, scale, ): pid_t = tl.program_id(0) pid_bh = tl.program_id(1).to(tl.int64) i_b = pid_bh // H i_h = pid_bh - i_b * H o_c = tl.arange(0, 64) o_k = tl.arange(0, 128) tok0 = pid_t * 64 rowb = (i_b * T * H + i_h) b_k = tl.load(k + rowb * 128 + (tok0 + o_c)[:, None] * (H * 128) + o_k[None, :]) b_g = tl.load(g + rowb * 128 + (tok0 + o_c)[:, None] * (H * 128) + o_k[None, :]) b_beta = tl.load(beta + rowb + (tok0 + o_c) * H).to(tl.float32) # cumsum via triangular-ones dot (tf32) instead of 6 shuffle rounds ltri = tl.where(o_c[:, None] >= o_c[None, :], 1.0, 0.0) b_g = tl.dot(ltri, b_g * RCP_LN2, input_precision="tf32") b_glast = tl.sum(tl.where((o_c[:, None] == 63), b_g, 0.0), 0) b_d = tl.exp2(b_glast) e_up = tl.exp2(b_g - b_glast[None, :]) e_all = e_up * b_d[None, :] b_kf = b_k.to(tl.float32) b_kl = (b_kf * e_up).to(tl.bfloat16) b_kr = (b_kf * tl.exp2(b_glast[None, :] - b_g)).to(tl.bfloat16) # ---- A0, B; M = (I+B)^-1 by bf16 doubling ---- b_B = tl.dot(b_kl, tl.trans(b_kr)) b_B = tl.where(o_c[:, None] > o_c[None, :], b_B * b_beta[:, None], 0.0).to(tl.bfloat16) eye = tl.where(o_c[:, None] == o_c[None, :], 1.0, 0.0).to(tl.bfloat16) X = (eye.to(tl.float32) - b_B.to(tl.float32)).to(tl.bfloat16) b_B = tl.dot(b_B, b_B).to(tl.bfloat16) # B^2 X = tl.dot(X, (eye.to(tl.float32) + b_B.to(tl.float32)).to(tl.bfloat16)).to(tl.bfloat16) b_B = tl.dot(b_B, b_B).to(tl.bfloat16) # B^4 X = tl.dot(X, (eye.to(tl.float32) + b_B.to(tl.float32)).to(tl.bfloat16)).to(tl.bfloat16) b_B = tl.dot(b_B, b_B).to(tl.bfloat16) # B^8 X = tl.dot(X, (eye.to(tl.float32) + b_B.to(tl.float32)).to(tl.bfloat16)).to(tl.bfloat16) b_B = tl.dot(b_B, b_B).to(tl.bfloat16) # B^16 X = tl.dot(X, (eye.to(tl.float32) + b_B.to(tl.float32)).to(tl.bfloat16)).to(tl.bfloat16) b_B = tl.dot(b_B, b_B).to(tl.bfloat16) # B^32 M = tl.dot(X, (eye.to(tl.float32) + b_B.to(tl.float32)).to(tl.bfloat16)).to(tl.bfloat16) # ---- w = M kb, u = M vb ---- b_v = tl.load(v + rowb * 128 + (tok0 + o_c)[:, None] * (H * 128) + o_k[None, :]) b_kb = (b_kf * (b_beta[:, None] * e_all)).to(tl.bfloat16) b_vb = (b_v.to(tl.float32) * b_beta[:, None]).to(tl.bfloat16) b_w = tl.dot(M, b_kb).to(tl.bfloat16) b_u = tl.dot(M, b_vb).to(tl.bfloat16) # ---- P'' = -kr^T w, C = kr^T u ---- pbase = (pid_bh * NT + pid_t) * (128 * 128) b_krt = tl.trans(b_kr) tl.store(P + pbase + o_k[:, None] * 128 + o_k[None, :], (-tl.dot(b_krt, b_w)).to(tl.bfloat16)) tl.store(C + pbase + o_k[:, None] * 128 + o_k[None, :], tl.dot(b_krt, b_u).to(tl.bfloat16)) tl.store(Pd + (pid_bh * NT + pid_t) * 128 + o_k, b_d) # ---- Aqk, qe, t ---- b_q = tl.load(q + rowb * 128 + (tok0 + o_c)[:, None] * (H * 128) + o_k[None, :]) b_ql = (b_q.to(tl.float32) * (scale * e_up)).to(tl.bfloat16) b_qg = (b_q.to(tl.float32) * (scale * e_all)).to(tl.bfloat16) b_Aqk = tl.dot(b_ql, tl.trans(b_kr)) b_Aqk = tl.where(o_c[:, None] >= o_c[None, :], b_Aqk, 0.0) b_Aqk = b_Aqk.to(tl.bfloat16) b_qe = (b_qg.to(tl.float32) - tl.dot(b_Aqk, b_w)).to(tl.bfloat16) b_t = tl.dot(b_Aqk, b_u).to(tl.bfloat16) tl.store(qe + rowb * 128 + (tok0 + o_c)[:, None] * (H * 128) + o_k[None, :], b_qe) tl.store(t + rowb * 128 + (tok0 + o_c)[:, None] * (H * 128) + o_k[None, :], b_t) # --------------------------------------------------------------------------- # K2 pass A: per-group local scan (S=0 init) + group composite A_g # grid (NV, BH, G), BV columns each; A maintained full-width redundantly # --------------------------------------------------------------------------- @triton.jit def _kda_k2a( P, C, Pd, A, Bg, T, H: tl.constexpr, NT, L: tl.constexpr, BV: tl.constexpr, ): i_v = tl.program_id(0) pid_bh = tl.program_id(1).to(tl.int64) i_g = tl.program_id(2) i_b = pid_bh // H i_h = pid_bh - i_b * H o_k = tl.arange(0, 128) o_v = i_v * BV + tl.arange(0, BV) rowb = (i_b * T * H + i_h) b_S = tl.zeros([128, BV], dtype=tl.float32) b_A = tl.where(o_k[:, None] == o_k[None, :], 1.0, 0.0) # [128,128] fp32 for j in range(L): i_t = i_g * L + j pbase = (pid_bh * NT + i_t) * (128 * 128) b_P = tl.load(P + pbase + o_k[:, None] * 128 + o_k[None, :]) b_C = tl.load(C + pbase + o_k[:, None] * 128 + o_v[None, :]) b_d = tl.load(Pd + (pid_bh * NT + i_t) * 128 + o_k) b_S = b_d[:, None] * b_S + b_C.to(tl.float32) + tl.dot(b_P, b_S.to(tl.bfloat16)) b_A = b_d[:, None] * b_A + tl.dot(b_P, b_A.to(tl.bfloat16)) abase = ((pid_bh * tl.num_programs(2)) + i_g) * (128 * 128) tl.store(A + abase + o_k[:, None] * 128 + o_k[None, :], b_A.to(tl.bfloat16)) tl.store(Bg + abase + o_k[:, None] * 128 + o_v[None, :], b_S.to(tl.bfloat16)) # --------------------------------------------------------------------------- # K2 pass B: tiny sequential scan over the G group composites -> boundary states # grid (NV, BH) # --------------------------------------------------------------------------- @triton.jit def _kda_k2b( A, Bg, Sb, H: tl.constexpr, G: tl.constexpr, BV: tl.constexpr, ): i_v = tl.program_id(0) pid_bh = tl.program_id(1).to(tl.int64) o_k = tl.arange(0, 128) o_v = i_v * BV + tl.arange(0, BV) b_S = tl.zeros([128, BV], dtype=tl.float32) for i_g in range(G): abase = ((pid_bh * G) + i_g) * (128 * 128) b_A = tl.load(A + abase + o_k[:, None] * 128 + o_k[None, :]) b_Bg = tl.load(Bg + abase + o_k[:, None] * 128 + o_v[None, :]) tl.store(Sb + abase + o_k[:, None] * 128 + o_v[None, :], b_S.to(tl.bfloat16)) b_S = b_Bg.to(tl.float32) + tl.dot(b_A, b_S.to(tl.bfloat16)) # --------------------------------------------------------------------------- # K2 pass C: final scan within each group from its boundary state + outputs # grid (NV, BH, G) # --------------------------------------------------------------------------- @triton.jit def _kda_k2c( P, C, Pd, qe, t, Sb, o, T, H: tl.constexpr, NT, L: tl.constexpr, BV: tl.constexpr, ): i_v = tl.program_id(0) pid_bh = tl.program_id(1).to(tl.int64) i_g = tl.program_id(2) i_b = pid_bh // H i_h = pid_bh - i_b * H o_c = tl.arange(0, 64) o_k = tl.arange(0, 128) o_v = i_v * BV + tl.arange(0, BV) rowb = (i_b * T * H + i_h) abase = ((pid_bh * tl.num_programs(2)) + i_g) * (128 * 128) b_S = tl.load(Sb + abase + o_k[:, None] * 128 + o_v[None, :]).to(tl.float32) for j in range(L): i_t = i_g * L + j tok0 = i_t * 64 pbase = (pid_bh * NT + i_t) * (128 * 128) b_P = tl.load(P + pbase + o_k[:, None] * 128 + o_k[None, :]) b_C = tl.load(C + pbase + o_k[:, None] * 128 + o_v[None, :]) b_d = tl.load(Pd + (pid_bh * NT + i_t) * 128 + o_k) b_qe = tl.load(qe + rowb * 128 + (tok0 + o_c)[:, None] * (H * 128) + o_k[None, :]) b_t = tl.load(t + rowb * 128 + (tok0 + o_c)[:, None] * (H * 128) + o_v[None, :]) b_o = tl.dot(b_qe, b_S.to(tl.bfloat16)) + b_t.to(tl.float32) tl.store(o + rowb * 128 + (tok0 + o_c)[:, None] * (H * 128) + o_v[None, :], b_o.to(tl.bfloat16)) b_S = b_d[:, None] * b_S + b_C.to(tl.float32) + tl.dot(b_P, b_S.to(tl.bfloat16)) # --------------------------------------------------------------------------- # K2 (fallback): single sequential scan with fused output # --------------------------------------------------------------------------- @triton.jit def _kda_k2( P, C, Pd, qe, t, o, T, H: tl.constexpr, NT, BV: tl.constexpr, ): i_v = tl.program_id(0) pid_bh = tl.program_id(1).to(tl.int64) i_b = pid_bh // H i_h = pid_bh - i_b * H o_c = tl.arange(0, 64) o_k = tl.arange(0, 128) o_v = i_v * BV + tl.arange(0, BV) rowb = (i_b * T * H + i_h) b_S = tl.zeros([128, BV], dtype=tl.float32) for i_t in range(NT): tok0 = i_t * 64 pbase = (pid_bh * NT + i_t) * (128 * 128) b_P = tl.load(P + pbase + o_k[:, None] * 128 + o_k[None, :]) # [128,128] b_C = tl.load(C + pbase + o_k[:, None] * 128 + o_v[None, :]) # [128,BV] b_d = tl.load(Pd + (pid_bh * NT + i_t) * 128 + o_k) # [128] fp32 b_qe = tl.load(qe + rowb * 128 + (tok0 + o_c)[:, None] * (H * 128) + o_k[None, :]) b_t = tl.load(t + rowb * 128 + (tok0 + o_c)[:, None] * (H * 128) + o_v[None, :]) b_o = tl.dot(b_qe, b_S.to(tl.bfloat16)) + b_t.to(tl.float32) # pre-update S tl.store(o + rowb * 128 + (tok0 + o_c)[:, None] * (H * 128) + o_v[None, :], b_o.to(tl.bfloat16)) b_S = b_d[:, None] * b_S + b_C.to(tl.float32) + tl.dot(b_P, b_S.to(tl.bfloat16)) # --------------------------------------------------------------------------- # Model # --------------------------------------------------------------------------- def _kda_forward(q, k, v, g, beta, scale, chunk_size=64): B, T, H, K = q.shape V = v.shape[-1] assert K == 128 and V == 128 and chunk_size == 64 and T % 64 == 0 NT = T // 64 BH = B * H dev = q.device bufs = _BUF_CACHE.get((B, T, H)) if bufs is None: # two-level scan split: aim for ~128 CTAs in k2a/k2c (2*BV-halves per bh) tgt = max(1, 128 // (2 * BH)) G = 1 while G * 2 <= tgt and NT % (G * 2) == 0: G *= 2 bufs = ( torch.empty(B, H, NT, K, V, dtype=torch.bfloat16, device=dev), # P torch.empty(B, H, NT, K, V, dtype=torch.bfloat16, device=dev), # C torch.empty(B * H * NT, K, dtype=torch.float32, device=dev), # Pd torch.empty(B, T, H, K, dtype=torch.bfloat16, device=dev), # qe torch.empty(B, T, H, V, dtype=torch.bfloat16, device=dev), # t torch.empty(B * H * G, K, V, dtype=torch.bfloat16, device=dev), # A torch.empty(B * H * G, K, V, dtype=torch.bfloat16, device=dev), # Bg torch.empty(B * H * G, K, V, dtype=torch.bfloat16, device=dev), # Sb G, ) _BUF_CACHE[(B, T, H)] = bufs P_, C_, Pd_, qe_, t_, A_, Bg_, Sb_, G = bufs L = NT // G o = torch.empty(B, T, H, V, dtype=torch.bfloat16, device=dev) _run_k1(q, k, v, g, beta, P_, C_, Pd_, qe_, t_, T, H, NT, scale, BH) _kda_k2a[(2, BH, G)](P_, C_, Pd_, A_, Bg_, T, H, NT, L=L, BV=64, num_warps=8, num_stages=2) _kda_k2b[(2, BH)](A_, Bg_, Sb_, H, G=G, BV=64, num_warps=8, num_stages=2) _kda_k2c[(2, BH, G)](P_, C_, Pd_, qe_, t_, Sb_, o, T, H, NT, L=L, BV=64, num_warps=8, num_stages=2) return o # ---------------- CUDA K1 (SM120 mma.sync path) ---------------- _CUDA_K1_SRC = r""" #include #include #define DEV __device__ __forceinline__ DEV int lane_id() { return threadIdx.x & 31; } DEV uint32_t smem_u32(const void* p) { return (uint32_t)__cvta_generic_to_shared(p); } DEV void ldm_x4(uint32_t& r0, uint32_t& r1, uint32_t& r2, uint32_t& r3, const void* p) { asm volatile("ldmatrix.sync.aligned.m8n8.x4.shared.b16 {%0,%1,%2,%3}, [%4];\n" : "=r"(r0), "=r"(r1), "=r"(r2), "=r"(r3) : "r"(smem_u32(p))); } DEV void ldm_x4t(uint32_t& r0, uint32_t& r1, uint32_t& r2, uint32_t& r3, const void* p) { asm volatile("ldmatrix.sync.aligned.m8n8.x4.trans.shared.b16 {%0,%1,%2,%3}, [%4];\n" : "=r"(r0), "=r"(r1), "=r"(r2), "=r"(r3) : "r"(smem_u32(p))); } DEV void stm_x4(const void* p, uint32_t r0, uint32_t r1, uint32_t r2, uint32_t r3) { asm volatile("stmatrix.sync.aligned.m8n8.x4.shared.b16 [%0], {%1,%2,%3,%4};\n" :: "r"(smem_u32(p)), "r"(r0), "r"(r1), "r"(r2), "r"(r3)); } DEV void mma_bf16(float& d0, float& d1, float& d2, float& d3, uint32_t a0, uint32_t a1, uint32_t a2, uint32_t a3, uint32_t b0, uint32_t b1) { asm volatile("mma.sync.aligned.m16n8k16.row.col.f32.bf16.bf16.f32 " "{%0,%1,%2,%3}, {%4,%5,%6,%7}, {%8,%9}, {%0,%1,%2,%3};\n" : "+f"(d0), "+f"(d1), "+f"(d2), "+f"(d3) : "r"(a0), "r"(a1), "r"(a2), "r"(a3), "r"(b0), "r"(b1)); } DEV uint32_t pack_bf16(float lo, float hi) { __nv_bfloat162 v = __floats2bfloat162_rn(lo, hi); return *reinterpret_cast(&v); } DEV float bf2f(__nv_bfloat16 x) { return __bfloat162float(x); } #ifdef K1_DBG __device__ bool k1_dbg_root() { return blockIdx.x == 0 && blockIdx.y == 0; } DEV void dump_f32(float* dbg, long base, const float* buf, int n) { if (blockIdx.x != 0 || blockIdx.y != 0) return; for (int i = threadIdx.x; i < n; i += 256) dbg[base + i] = buf[i]; } DEV void dump_bf(float* dbg, long base, const __nv_bfloat16* buf, int n) { if (!k1_dbg_root()) return; for (int i = threadIdx.x; i < n; i += 256) dbg[base + i] = bf2f(buf[i]); } #endif // row pitch must be a multiple of 8 elems (16B units): S % 8 == 0 makes the XOR // swizzle conflict-free for ldmatrix/stmatrix 8-row groups (unit%S differs only via q^r). constexpr int LD128 = 128; constexpr int LD64 = 64; constexpr float RCP_LN2 = 1.4426950408889634f; struct Smem { __nv_bfloat16 kl[64 * LD128]; // k*e_up (A0 operand) -> qe/t/P/C store staging __nv_bfloat16 qlq[64 * LD128]; // q*scale*e_up (Aqk operand), written in prologue __nv_bfloat16 krt[128 * LD64]; // kr^T row-major [k=128][c=64] __nv_bfloat16 kb[64 * LD128]; // k*beta*e_all -> w __nv_bfloat16 vb[64 * LD128]; // v*beta -> u __nv_bfloat16 ba[64 * LD64]; // chain ping / M __nv_bfloat16 bb[64 * LD64]; // chain pong / Aqk float beta_s[64]; float d_s[128]; float part[512]; }; // swizzled offsets: XOR the 8-element chunk index with (row & 7) DEV int o128(int r, int c) { return r * LD128 + ((c >> 3) ^ (r & 7)) * 8 + (c & 7); } DEV int o64(int r, int c) { return r * LD64 + ((c >> 3) ^ (r & 7)) * 8 + (c & 7); } // acc[RM/16][RN/8][4] += A @ B for this warp's [RM, RN] slice. // A: row-major [M,K] smem (swizzled), warp rows arow0..+RM, row pitch aw // B: row-major [K,N] smem (swizzled), warp cols bcol0..+RN, row pitch bw template DEV void dot_mn(float acc[RM / 16][RN / 8][4], const __nv_bfloat16* A, int arow0, int aw, const __nv_bfloat16* B, int bcol0, int bw, int K) { int l = lane_id(); for (int k0 = 0; k0 < K; k0 += 16) { uint32_t a[RM / 16][4]; #pragma unroll for (int mi = 0; mi < RM / 16; mi++) { int r = arow0 + mi * 16 + (l % 16); int c = k0 + (l / 16) * 8; ldm_x4(a[mi][0], a[mi][1], a[mi][2], a[mi][3], A + (aw == LD128 ? o128(r, c) : o64(r, c))); } uint32_t b[RN / 4]; // each x4t yields 2 frags x 2 regs #pragma unroll for (int np = 0; np < RN / 16; np++) { uint32_t r0, r1, r2, r3; int g = l / 8; int kk = k0 + (l % 8) + (g & 1) * 8; int nn = bcol0 + np * 16 + (g >> 1) * 8; ldm_x4t(r0, r1, r2, r3, B + (bw == LD128 ? o128(kk, nn) : o64(kk, nn))); b[np * 4 + 0] = r0; b[np * 4 + 1] = r1; b[np * 4 + 2] = r2; b[np * 4 + 3] = r3; } #pragma unroll for (int mi = 0; mi < RM / 16; mi++) #pragma unroll for (int ni = 0; ni < RN / 8; ni++) mma_bf16(acc[mi][ni][0], acc[mi][ni][1], acc[mi][ni][2], acc[mi][ni][3], a[mi][0], a[mi][1], a[mi][2], a[mi][3], b[ni * 2], b[ni * 2 + 1]); } } // zero acc template DEV void zero_acc(float acc[RM / 16][RN / 8][4]) { #pragma unroll for (int i = 0; i < RM / 16; i++) #pragma unroll for (int j = 0; j < RN / 8; j++) #pragma unroll for (int p = 0; p < 4; p++) acc[i][j][p] = 0.f; } // stage warp's [RM, RN] acc slice to a 64- or 128-wide smem tile at (row0, col0). // mode: 0 plain, 1 negate. LD picks o64/o128 addressing. template DEV void stage_acc(float acc[RM / 16][RN / 8][4], int row0, int col0, __nv_bfloat16* base, bool negate) { int l = lane_id(); #pragma unroll for (int mi = 0; mi < RM / 16; mi++) #pragma unroll for (int ni = 0; ni < RN / 8; ni += 2) { int r = row0 + mi * 16; int c = col0 + ni * 8; int lr = r + (l % 8) + ((l / 8) & 1) * 8; int lc = c + (l / 16) * 8; int off = (LD == LD128) ? o128(lr, lc) : o64(lr, lc); float s0 = negate ? -acc[mi][ni][0] : acc[mi][ni][0]; float s1 = negate ? -acc[mi][ni][1] : acc[mi][ni][1]; float s2 = negate ? -acc[mi][ni][2] : acc[mi][ni][2]; float s3 = negate ? -acc[mi][ni][3] : acc[mi][ni][3]; float u0 = negate ? -acc[mi][ni + 1][0] : acc[mi][ni + 1][0]; float u1 = negate ? -acc[mi][ni + 1][1] : acc[mi][ni + 1][1]; float u2 = negate ? -acc[mi][ni + 1][2] : acc[mi][ni + 1][2]; float u3 = negate ? -acc[mi][ni + 1][3] : acc[mi][ni + 1][3]; stm_x4(base + off, pack_bf16(s0, s1), pack_bf16(s2, s3), pack_bf16(u0, u1), pack_bf16(u2, u3)); } } // copy a [64,128] bf16 smem tile (swizzled, pitch LD128) to gmem rows rowG..rowG+63, // gmem row pitch = gstride elements, i.e. dst + (rowG+i)*gstride + j. DEV void tile_to_gmem(const __nv_bfloat16* src, __nv_bfloat16* dst, int rowG, long gstride) { int t = threadIdx.x; int r = t >> 2; int u = (t & 3) * 4; // uint4 index (8 bf16 each) const uint4* s4 = reinterpret_cast(src); uint4* d4 = reinterpret_cast(dst); long grow = (long)(rowG + r) * (gstride >> 3); // gstride in bf16 -> in uint4 #pragma unroll for (int i = 0; i < 4; i++) { int uu = u + i; int soff = o128(r, uu * 8) >> 3; // /8 elems -> /8 in uint4? (8 bf16 = 1 uint4) but swizzle offset is in elems: elem offset o128(r, uu*8) is a multiple of 8 -> uint4 index = /8 d4[grow + uu] = s4[soff]; } } // --------------------------------------------------------------------------- // K1: grid (NT, BH), 256 threads (8 warps: wm = w>>2, wn = w&3). // --------------------------------------------------------------------------- __global__ void __launch_bounds__(256, 1) kda_k1( const __nv_bfloat16* __restrict__ q, const __nv_bfloat16* __restrict__ k, const __nv_bfloat16* __restrict__ v, const float* __restrict__ g, const __nv_bfloat16* __restrict__ beta, __nv_bfloat16* __restrict__ P, __nv_bfloat16* __restrict__ Cm, float* __restrict__ Pd, __nv_bfloat16* __restrict__ qe, __nv_bfloat16* __restrict__ tt, long T, int H, long NT, float scale #ifdef K1_DBG , float* dbg #endif ) { extern __shared__ __nv_bfloat16 smem_raw[]; Smem& s = *reinterpret_cast(smem_raw); int pid_t = blockIdx.x; long bh = blockIdx.y; long i_b = bh / H, i_h = bh - i_b * H; long tok0 = (long)pid_t * 64; int tid = threadIdx.x, l = lane_id(), w = tid >> 5; int wm = w >> 2, wn = w & 3; long HK = (long)H * 128; long rowK = (i_b * T + tok0) * HK + i_h * 128; // k/v/q/g row (t-major, K-contig) long rowB = (i_b * T + tok0) * H + i_h; // beta row long outQ = (i_b * T + tok0) * HK + i_h * 128; // qe/t rows // ---------------- prologue ---------------- // thread owns 4 columns c0..c0+3 and 8 rows r0..r0+7 (rg = tid>>5, cp = tid&31). // All gmem loads vectorized: g float4, k/v/q 8B. part[] overlays s.ba (free here). { int cp = tid & 31; int c0 = cp * 4; int rg = tid >> 5; int r0 = rg * 8; const float4* gp = reinterpret_cast(g + rowK + (long)r0 * HK + c0); const uint2* kp = reinterpret_cast(k + rowK + (long)r0 * HK + c0); const uint2* vp = reinterpret_cast(v + rowK + (long)r0 * HK + c0); const uint2* qp = reinterpret_cast(q + rowK + (long)r0 * HK + c0); float4 gx[8]; uint2 kr4[8], vr4[8], qr4[8]; #pragma unroll for (int i = 0; i < 8; i++) gx[i] = gp[(long)i * (HK >> 2)]; #pragma unroll for (int i = 0; i < 8; i++) kr4[i] = kp[(long)i * (HK >> 2)]; #pragma unroll for (int i = 0; i < 8; i++) vr4[i] = vp[(long)i * (HK >> 2)]; #pragma unroll for (int i = 0; i < 8; i++) qr4[i] = qp[(long)i * (HK >> 2)]; float acc[8][4]; float run[4] = {0.f, 0.f, 0.f, 0.f}; #pragma unroll for (int i = 0; i < 8; i++) { run[0] += gx[i].x * RCP_LN2; acc[i][0] = run[0]; run[1] += gx[i].y * RCP_LN2; acc[i][1] = run[1]; run[2] += gx[i].z * RCP_LN2; acc[i][2] = run[2]; run[3] += gx[i].w * RCP_LN2; acc[i][3] = run[3]; } float* part = reinterpret_cast(s.ba); // [8][128], ba not used yet #pragma unroll for (int j = 0; j < 4; j++) part[rg * 128 + c0 + j] = run[j]; if (tid < 64) s.beta_s[tid] = bf2f(beta[rowB + (long)tid * H]); __syncthreads(); float glast[4], pbelow[4]; #pragma unroll for (int j = 0; j < 4; j++) { float tot = 0.f, below = 0.f; for (int rr = 0; rr < 8; rr++) { float p = part[rr * 128 + c0 + j]; tot += p; if (rr < rg) below += p; } glast[j] = tot; pbelow[j] = below; } if (rg == 0) { float4 dd = make_float4(exp2f(glast[0]), exp2f(glast[1]), exp2f(glast[2]), exp2f(glast[3])); reinterpret_cast(s.d_s)[cp] = dd; reinterpret_cast(Pd + (bh * NT + pid_t) * 128)[cp] = dd; } __nv_bfloat162* kl2 = reinterpret_cast<__nv_bfloat162*>(s.kl); __nv_bfloat162* ql2 = reinterpret_cast<__nv_bfloat162*>(s.qlq); __nv_bfloat162* kb2 = reinterpret_cast<__nv_bfloat162*>(s.kb); __nv_bfloat162* vb2 = reinterpret_cast<__nv_bfloat162*>(s.vb); #pragma unroll for (int i = 0; i < 8; i++) { int r = r0 + i; __nv_bfloat162 k01 = *reinterpret_cast<__nv_bfloat162*>(&kr4[i]); __nv_bfloat162 k23 = *reinterpret_cast<__nv_bfloat162*>(&kr4[i].y); __nv_bfloat162 v01 = *reinterpret_cast<__nv_bfloat162*>(&vr4[i]); __nv_bfloat162 v23 = *reinterpret_cast<__nv_bfloat162*>(&vr4[i].y); __nv_bfloat162 q01 = *reinterpret_cast<__nv_bfloat162*>(&qr4[i]); __nv_bfloat162 q23 = *reinterpret_cast<__nv_bfloat162*>(&qr4[i].y); float bt = s.beta_s[r]; float e_up[4], e_dn[4], ea[4]; #pragma unroll for (int j = 0; j < 4; j++) { float gc = acc[i][j] + pbelow[j]; e_up[j] = exp2f(gc - glast[j]); e_dn[j] = exp2f(glast[j] - gc); ea[j] = exp2f(gc); } int o = o128(r, c0) >> 1; kl2[o] = __floats2bfloat162_rn(bf2f(k01.x) * e_up[0], bf2f(k01.y) * e_up[1]); kl2[o + 1] = __floats2bfloat162_rn(bf2f(k23.x) * e_up[2], bf2f(k23.y) * e_up[3]); ql2[o] = __floats2bfloat162_rn(bf2f(q01.x) * scale * e_up[0], bf2f(q01.y) * scale * e_up[1]); ql2[o + 1] = __floats2bfloat162_rn(bf2f(q23.x) * scale * e_up[2], bf2f(q23.y) * scale * e_up[3]); kb2[o] = __floats2bfloat162_rn(bf2f(k01.x) * bt * ea[0], bf2f(k01.y) * bt * ea[1]); kb2[o + 1] = __floats2bfloat162_rn(bf2f(k23.x) * bt * ea[2], bf2f(k23.y) * bt * ea[3]); vb2[o] = __floats2bfloat162_rn(bf2f(v01.x) * bt, bf2f(v01.y) * bt); vb2[o + 1] = __floats2bfloat162_rn(bf2f(v23.x) * bt, bf2f(v23.y) * bt); s.krt[o64(c0, r)] = __float2bfloat16(bf2f(k01.x) * e_dn[0]); s.krt[o64(c0 + 1, r)] = __float2bfloat16(bf2f(k01.y) * e_dn[1]); s.krt[o64(c0 + 2, r)] = __float2bfloat16(bf2f(k23.x) * e_dn[2]); s.krt[o64(c0 + 3, r)] = __float2bfloat16(bf2f(k23.y) * e_dn[3]); } __syncthreads(); } #ifdef K1_DBG dump_bf(dbg, 0, s.kl, 64*LD128); dump_bf(dbg, 61440, s.qlq, 64*LD128); dump_f32(dbg, 70000, s.beta_s, 64); dump_bf(dbg, 8192, s.krt, 128*LD64); dump_bf(dbg, 16384, s.kb, 64*LD128); dump_bf(dbg, 24576, s.vb, 64*LD128); __syncthreads(); #endif // ---------------- A0 = kl @ krT ; B = tril(A0*beta, -1) ; X = I - B ---------------- #ifdef K1_EARLY if (K1_EARLY == 1) return; #endif float a0[2][2][4]; zero_acc<32, 16>(a0); dot_mn<32, 16>(a0, s.kl, wm * 32, LD128, s.krt, wn * 16, LD64, 128); __syncthreads(); // kl free after all warps' A0 #pragma unroll for (int mi = 0; mi < 2; mi++) #pragma unroll for (int ni = 0; ni < 2; ni++) #pragma unroll for (int p = 0; p < 4; p++) { int r, c; // frag coords within [32,16] warp slice r = wm * 32 + mi * 16 + (l >> 2); c = wn * 16 + ni * 8 + ((l & 3) * 2 + (p & 1)); if (p >= 2) r += 8; float val = (r > c) ? a0[mi][ni][p] * s.beta_s[r] : 0.f; a0[mi][ni][p] = val; } float xacc[2][2][4]; #pragma unroll for (int mi = 0; mi < 2; mi++) #pragma unroll for (int ni = 0; ni < 2; ni++) #pragma unroll for (int p = 0; p < 4; p++) { int r = wm * 32 + mi * 16 + (l >> 2) + (p >= 2 ? 8 : 0); int c = wn * 16 + ni * 8 + ((l & 3) * 2 + (p & 1)); xacc[mi][ni][p] = (r == c) ? 1.f : -a0[mi][ni][p]; } #ifdef K1_DBG if (blockIdx.x == 0 && blockIdx.y == 0 && w == 0) for (int mi = 0; mi < 2; mi++) for (int ni = 0; ni < 2; ni++) for (int p = 0; p < 4; p++) dbg[71000 + mi * 8 + ni * 4 + p] = a0[mi][ni][p]; __syncthreads(); #endif stage_acc<32, 16, LD64>(a0, wm * 32, wn * 16, s.ba, false); __syncthreads(); #ifdef K1_DBG dump_bf(dbg, 32768, s.ba, 64*LD64); __syncthreads(); #endif #ifdef K1_EARLY if (K1_EARLY == 2) return; #endif // ---------------- chain: X <- X + X @ B^(2^i) ---------------- __nv_bfloat16* cur = s.ba; __nv_bfloat16* oth = s.bb; #pragma unroll for (int it = 0; it < 5; it++) { float sq[2][2][4]; zero_acc<32, 16>(sq); dot_mn<32, 16>(sq, cur, wm * 32, LD64, cur, wn * 16, LD64, 64); stage_acc<32, 16, LD64>(sq, wm * 32, wn * 16, oth, false); __syncthreads(); // cur dead: all warps done squaring stage_acc<32, 16, LD64>(xacc, wm * 32, wn * 16, cur, false); __syncthreads(); dot_mn<32, 16>(xacc, cur, wm * 32, LD64, oth, wn * 16, LD64, 64); // acc-init = X __syncthreads(); __nv_bfloat16* tmp = cur; cur = oth; oth = tmp; } // M = xacc stage_acc<32, 16, LD64>(xacc, wm * 32, wn * 16, s.ba, false); __syncwarp(); #ifdef K1_DBG dump_bf(dbg, 36864, s.ba, 64*LD64); __syncthreads(); #endif // ---------------- w = M @ kb ; u = M @ vb ---------------- float wacc[2][4][4], uacc[2][4][4]; zero_acc<32, 32>(wacc); zero_acc<32, 32>(uacc); dot_mn<32, 32>(wacc, s.ba, wm * 32, LD64, s.kb, wn * 32, LD128, 64); dot_mn<32, 32>(uacc, s.ba, wm * 32, LD64, s.vb, wn * 32, LD128, 64); stage_acc<32, 32, LD128>(wacc, wm * 32, wn * 32, s.kb, false); stage_acc<32, 32, LD128>(uacc, wm * 32, wn * 32, s.vb, false); __syncthreads(); #ifdef K1_DBG dump_bf(dbg, 40960, s.kb, 64*LD128); dump_bf(dbg, 49152, s.vb, 64*LD128); __syncthreads(); #endif #ifdef K1_EARLY if (K1_EARLY == 3) return; #endif // ---------------- Aqk = tril(ql @ krT, incl diag) ; ql precomputed in prologue ---------------- float aqk[2][2][4]; zero_acc<32, 16>(aqk); dot_mn<32, 16>(aqk, s.qlq, wm * 32, LD128, s.krt, wn * 16, LD64, 128); #pragma unroll for (int mi = 0; mi < 2; mi++) #pragma unroll for (int ni = 0; ni < 2; ni++) #pragma unroll for (int p = 0; p < 4; p++) { int r = wm * 32 + mi * 16 + (l >> 2) + (p >= 2 ? 8 : 0); int c = wn * 16 + ni * 8 + ((l & 3) * 2 + (p & 1)); aqk[mi][ni][p] = (r >= c) ? aqk[mi][ni][p] : 0.f; // ql already has scale } stage_acc<32, 16, LD64>(aqk, wm * 32, wn * 16, s.bb, false); __syncthreads(); #ifdef K1_DBG dump_bf(dbg, 57344, s.bb, 64*LD64); #endif // ---------------- qe = ql*d - Aqk@w ; t = Aqk@u ---------------- float qeacc[2][4][4], tacc[2][4][4]; zero_acc<32, 32>(qeacc); zero_acc<32, 32>(tacc); dot_mn<32, 32>(qeacc, s.bb, wm * 32, LD64, s.kb, wn * 32, LD128, 64); dot_mn<32, 32>(tacc, s.bb, wm * 32, LD64, s.vb, wn * 32, LD128, 64); // pack qe with + ql[r][c]*d[c], stage into kl (element-wise read then overwrite) #pragma unroll for (int mi = 0; mi < 2; mi++) #pragma unroll for (int ni = 0; ni < 4; ni += 2) { int r = wm * 32 + mi * 16; int cc = wn * 32 + ni * 8; int lr = r + (l % 8) + ((l / 8) & 1) * 8; int lc = cc + (l / 16) * 8; int off = o128(lr, lc); int rr = r + (l >> 2); // fragment row (address row differs!) int cb = cc + 2 * (l & 3); // fragment cols float q00 = bf2f(s.qlq[o128(rr, cb)]) * s.d_s[cb]; float q01 = bf2f(s.qlq[o128(rr, cb + 1)]) * s.d_s[cb + 1]; float q02 = bf2f(s.qlq[o128(rr + 8, cb)]) * s.d_s[cb]; float q03 = bf2f(s.qlq[o128(rr + 8, cb + 1)]) * s.d_s[cb + 1]; float q10 = bf2f(s.qlq[o128(rr, cb + 8)]) * s.d_s[cb + 8]; float q11 = bf2f(s.qlq[o128(rr, cb + 9)]) * s.d_s[cb + 9]; float q12 = bf2f(s.qlq[o128(rr + 8, cb + 8)]) * s.d_s[cb + 8]; float q13 = bf2f(s.qlq[o128(rr + 8, cb + 9)]) * s.d_s[cb + 9]; stm_x4(s.kl + off, pack_bf16(q00 - qeacc[mi][ni][0], q01 - qeacc[mi][ni][1]), pack_bf16(q02 - qeacc[mi][ni][2], q03 - qeacc[mi][ni][3]), pack_bf16(q10 - qeacc[mi][ni + 1][0], q11 - qeacc[mi][ni + 1][1]), pack_bf16(q12 - qeacc[mi][ni + 1][2], q13 - qeacc[mi][ni + 1][3])); } __syncwarp(); __syncthreads(); tile_to_gmem(s.kl, qe + outQ, 0, HK); __syncthreads(); stage_acc<32, 32, LD128>(tacc, wm * 32, wn * 32, s.kl, false); __syncthreads(); tile_to_gmem(s.kl, tt + outQ, 0, HK); __syncthreads(); #ifdef K1_EARLY if (K1_EARLY == 4) return; #endif // ---------------- P = -(krt @ w) ; C = krt @ u (all 8 warps, [64,32] tiles) -------- // output rows 0-63 stage into qlq, rows 64-127 into kl (both free); kb/vb hold w/u. long pbase = ((i_b * H + i_h) * NT + pid_t) * 128 * 128; { float pacc[4][4][4]; zero_acc<64, 32>(pacc); dot_mn<64, 32>(pacc, s.krt, wm * 64, LD64, s.kb, wn * 32, LD128, 64); __syncthreads(); stage_acc<64, 32, LD128>(pacc, 0, wn * 32, wm ? s.kl : s.qlq, true); __syncthreads(); tile_to_gmem(s.qlq, P + pbase, 0, 128); tile_to_gmem(s.kl, P + pbase + 64 * 128, 0, 128); __syncthreads(); } { float cacc[4][4][4]; zero_acc<64, 32>(cacc); dot_mn<64, 32>(cacc, s.krt, wm * 64, LD64, s.vb, wn * 32, LD128, 64); __syncthreads(); stage_acc<64, 32, LD128>(cacc, 0, wn * 32, wm ? s.kl : s.qlq, false); __syncthreads(); tile_to_gmem(s.qlq, Cm + pbase, 0, 128); tile_to_gmem(s.kl, Cm + pbase + 64 * 128, 0, 128); __syncthreads(); } } void kda_k1_launch(torch::Tensor q, torch::Tensor k, torch::Tensor v, torch::Tensor g, torch::Tensor beta, torch::Tensor P, torch::Tensor C, torch::Tensor Pd, torch::Tensor qe, torch::Tensor t, long T, long H, long NT, double scale #ifdef K1_DBG , torch::Tensor dbg #endif ) { dim3 grid((unsigned)NT, (unsigned)(q.size(0) * H)); size_t smem = sizeof(Smem); static bool set = false; if (!set) { cudaError_t e = cudaFuncSetAttribute(kda_k1, cudaFuncAttributeMaxDynamicSharedMemorySize, (int)smem); TORCH_CHECK(e == cudaSuccess, "setattr failed: ", cudaGetErrorString(e), " smem=", smem); set = true; } cudaError_t el = cudaGetLastError(); TORCH_CHECK(el == cudaSuccess, "pre-launch: ", cudaGetErrorString(el)); kda_k1<<>>( reinterpret_cast(q.data_ptr()), reinterpret_cast(k.data_ptr()), reinterpret_cast(v.data_ptr()), g.data_ptr(), reinterpret_cast(beta.data_ptr()), reinterpret_cast<__nv_bfloat16*>(P.data_ptr()), reinterpret_cast<__nv_bfloat16*>(C.data_ptr()), Pd.data_ptr(), reinterpret_cast<__nv_bfloat16*>(qe.data_ptr()), reinterpret_cast<__nv_bfloat16*>(t.data_ptr()), T, (int)H, NT, (float)scale #ifdef K1_DBG , dbg.data_ptr() #endif ); } """ def _build_cuda_k1(): """Compile the custom CUDA K1 kernel; returns None if unavailable.""" try: if not torch.cuda.is_available(): return None major, _ = torch.cuda.get_device_capability(0) if major != 12: return None from torch.utils.cpp_extension import load_inline bdir = os.path.join(os.path.dirname(os.path.abspath(__file__)), "_build") os.makedirs(bdir, exist_ok=True) return load_inline( name="kda_cuda_k1_v3", cpp_sources="void kda_k1_launch(torch::Tensor q, torch::Tensor k, torch::Tensor v, " "torch::Tensor g, torch::Tensor beta, torch::Tensor P, torch::Tensor C, " "torch::Tensor Pd, torch::Tensor qe, torch::Tensor t, long T, long H, " "long NT, double scale);", cuda_sources=_CUDA_K1_SRC, functions=["kda_k1_launch"], extra_cuda_cflags=["-O3", "--use_fast_math", "-arch=sm_120a", "-std=c++20"], build_directory=bdir, verbose=False, ) except Exception: return None _CUDA_K1 = None _CUDA_K1_TRIED = False def _run_k1(q, k, v, g, beta, P_, C_, Pd_, qe_, t_, T, H, NT, scale, BH=1): """Dispatch K1 to the CUDA kernel when available, else Triton.""" global _CUDA_K1, _CUDA_K1_TRIED if not _CUDA_K1_TRIED: _CUDA_K1 = _build_cuda_k1() _CUDA_K1_TRIED = True if _CUDA_K1 is not None: _CUDA_K1.kda_k1_launch(q, k, v, g, beta, P_, C_, Pd_, qe_, t_, T, H, NT, scale) else: _kda_k1[(NT, BH_global(q))]( q, k, v, g, beta, P_, C_, Pd_, qe_, t_, T, H, NT, scale, num_warps=8) _BUF_CACHE: dict = {} 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_forward(q, k, v, g, beta, self.scale, 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]