code wiki / _hdl_build / nx_nofloat_kda.nx

nx_nofloat_kda.nx source

↩ module page · 62 lines · 3482 B

1// nx_nofloat_kda.nx -- KIMI DELTA ATTENTION (KDA) in sovereign no-float = the K3 namesake attention 2// (operator 2026-07-19 "state of the art... clearly researched, novel, evidence driven"). Extends the proven 3// nx_nofloat_linattn recurrent kernel with the TWO mechanisms that make KDA distinct (arXiv 2510.26692, Kimi 4// Linear; 3:1 KDA:MLA in K3): 5// (a) DELTA RULE (error-correcting write): instead of superposing, KDA overwrites the memory along k -- 6// pred_t = readout(S_{t-1}, k_t) = k_t^T S_{t-1} ; delta_t = beta_t (v_t - pred_t) 7// S_t = S_{t-1} + k_t (x) delta_t (writing the SAME key twice OVERWRITES, not doubles) 8// (b) FINE-GRAINED (per-channel diagonal) GATING: each key-channel forgets at its OWN learned rate -- 9// S_t[i][j] = alpha_i * S_{t-1}[i][j] (alpha per key-dim i; GLA has ONE scalar, KDA has d gates) 10// Combined step: S_t = Diag(alpha) S_{t-1} then delta-write(k_t, v_t, beta_t) ; o_t = q_t S_t. 11// Pure integer Q16 -> BIT-EXACT deterministic (a float delta/gated stack drifts by accumulation order; ours 12// does not). Reuses la_readout/la_accum/la_zero (DRY, rule-15). Scratch allocated ONCE per forward (no 13// per-step mmap -> no leak, the mkv2fmp4 lesson). license_tier: ORIGINAL No hw writes (Rule 26). 14import "nx_nofloat_linattn.nx" 15import "nx_syscalls.nx" 16 17const KDA_QBITS: i64 = 16 18const KDA_Q: i64 = 65536 19 20// per-channel diagonal gating: S[i][j] <- (alpha[i] * S[i][j]) >> 16 (each key-row decays at its own rate) 21func kda_gate(s: *i64, alpha: *i64, d: i64) -> i64 { 22 var i: i64 = 0 23 while i < d { 24 let ai: i64 = alpha[i] 25 var j: i64 = 0 26 while j < d { s[i*d+j] = (ai * s[i*d+j]) >> KDA_QBITS; j = j + 1 } 27 i = i + 1 28 } 29 return 0 30} 31// delta-rule write: pred = readout(S,k); delta[j] = (beta*(v[j]-pred[j]))>>16; S += k (x) delta. 32// pred and delta are caller-owned scratch (length d) -- allocated once per forward, never per step. 33func kda_write(s: *i64, k: *i64, v: *i64, beta: i64, d: i64, pred: *i64, delta: *i64) -> i64 { 34 la_readout(s, k, pred, d) // pred = k^T S (current stored value along k) 35 var j: i64 = 0 36 while j < d { delta[j] = (beta * (v[j] - pred[j])) >> KDA_QBITS; j = j + 1 } 37 la_accum(s, k, delta, d) // S += k (x) delta (error-correcting) 38 return 0 39} 40// KDA forward over a length-T sequence. Q,K,V are T x d (Q16). alpha = length-d per-channel gate (Q16, 41// <=65536). beta = scalar write strength (Q16). Writes O (T x d). Gate-then-delta-write-then-read per step. 42func kda_forward(qm: *i64, km: *i64, vm: *i64, om: *i64, alpha: *i64, beta: i64, t: i64, d: i64) -> i64 { 43 let s: *i64 = sys_mmap(LA_DMAX*LA_DMAX*8) as *i64 44 let pred: *i64 = sys_mmap(LA_DMAX*8) as *i64 45 let delta: *i64 = sys_mmap(LA_DMAX*8) as *i64 46 la_zero(s, d) 47 var ti: i64 = 0 48 while ti < t { 49 kda_gate(s, alpha, d) 50 kda_write(s, ((km as i64) + ti*d*8) as *i64, ((vm as i64) + ti*d*8) as *i64, beta, d, pred, delta) 51 la_readout(s, ((qm as i64) + ti*d*8) as *i64, ((om as i64) + ti*d*8) as *i64, d) 52 ti = ti + 1 53 } 54 return 0 55} 56// exposed single-step primitives for the gate: gated delta-write into a live state, then query. 57func kda_step(s: *i64, k: *i64, v: *i64, q: *i64, o: *i64, alpha: *i64, beta: i64, d: i64, pred: *i64, delta: *i64) -> i64 { 58 kda_gate(s, alpha, d) 59 kda_write(s, k, v, beta, d, pred, delta) 60 la_readout(s, q, o, d) 61 return 0 62}