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}