code wiki / (root) / nx_kda_gate.nx

nx_kda_gate.nx source

↩ module page · 171 lines · 7108 B

1// nx_kda_gate.nx -- two things nx_nofloat_kda_gate does NOT cover: 2// (A) the ALGEBRAIC IDENTITY that licenses the shipped kernel's factored update, and 3// (B) the KDA:MLA interleave schedule, a separate census feature that no organ implemented. 4// 5// (A) WHY IT MATTERS. Kimi Linear (arXiv 2510.26692) states the rule as 6// S_t = (I - beta k k^T) Diag(alpha) S_{t-1} + beta k v^T 7// which literally builds a d_k x d_k matrix: O(d_k^2 * d_v). The shipped kernel 8// (_hdl_build/nx_nofloat_kda.nx) does NOT do that -- it computes pred = k^T S, then S += k (x) beta(v-pred), 9// which is O(d_k * d_v). That is a real optimisation and it is CORRECT only because (I - beta k k^T) is 10// identity-plus-rank-1. nx_nofloat_kda_gate proves the kernel's BEHAVIOUR (overwrite, gating, determinism); 11// nothing proved the factorisation itself. T1 does, over two steps from a NON-ZERO state -- one step from 12// S=0 would agree even for a wrong factorisation. 13// 14// T3 guards against proving the identity on a DEGENERATE case: if k were a basis vector, (I - beta k k^T) 15// would be diagonal and both forms would agree trivially. T3 asserts the off-diagonal term is non-zero, so 16// T1 exercised the genuinely non-diagonal path. 17// 18// Deliberately NOT retested here (owned by nx_nofloat_kda_gate 5/5): delta-overwrite idempotence, 19// per-channel vs scalar gating, bit-exact determinism. 20// license_tier: ORIGINAL No hw writes (Rule 26). expect_exit: 0 21import "nx_syscalls.nx" 22import "nx_gate_verdict.nx" 23import "nx_kda.nx" 24 25const KG_FP: i64 = 1024 26 27func kg_mul(a: i64, b: i64) -> i64 { return (a * b) / KG_FP } 28 29func kg_state(dk: i64, dv: i64) -> *i64 { 30 let S: *i64 = sys_mmap(dk * dv * 8) as *i64 31 var i: i64 = 0 32 while i < dk * dv { S[i] = 0; i = i + 1 } 33 return S 34} 35 36// the FACTORED update the shipped kernel uses: pred = k^T Diag(alpha) S ; S = Diag(alpha) S + beta k (v-pred)^T 37func kg_step_factored(S: *i64, dk: i64, dv: i64, alpha: *i64, beta: i64, k: *i64, v: *i64, sc: *i64) -> i64 { 38 var i: i64 = 0 39 while i < dk { 40 var j: i64 = 0 41 while j < dv { S[i * dv + j] = kg_mul(alpha[i], S[i * dv + j]); j = j + 1 } 42 i = i + 1 43 } 44 var j2: i64 = 0 45 while j2 < dv { sc[j2] = 0; j2 = j2 + 1 } 46 i = 0 47 while i < dk { 48 var j: i64 = 0 49 while j < dv { sc[j] = sc[j] + kg_mul(k[i], S[i * dv + j]); j = j + 1 } 50 i = i + 1 51 } 52 j2 = 0 53 while j2 < dv { sc[j2] = v[j2] - sc[j2]; j2 = j2 + 1 } 54 i = 0 55 while i < dk { 56 let bk: i64 = kg_mul(beta, k[i]) 57 var j: i64 = 0 58 while j < dv { S[i * dv + j] = S[i * dv + j] + kg_mul(bk, sc[j]); j = j + 1 } 59 i = i + 1 60 } 61 return 0 62} 63 64// the LITERAL rule, materialising (I - beta k k^T). Reference only. 65func kg_step_naive(S: *i64, dk: i64, dv: i64, alpha: *i64, beta: i64, k: *i64, v: *i64) -> i64 { 66 let Sd: *i64 = sys_mmap(dk * dv * 8) as *i64 67 var i: i64 = 0 68 while i < dk { 69 var j: i64 = 0 70 while j < dv { Sd[i * dv + j] = kg_mul(alpha[i], S[i * dv + j]); j = j + 1 } 71 i = i + 1 72 } 73 let M: *i64 = sys_mmap(dk * dk * 8) as *i64 74 i = 0 75 while i < dk { 76 var m: i64 = 0 77 while m < dk { 78 var e: i64 = 0 79 if i == m { e = KG_FP } 80 M[i * dk + m] = e - kg_mul(kg_mul(beta, k[i]), k[m]) 81 m = m + 1 82 } 83 i = i + 1 84 } 85 i = 0 86 while i < dk { 87 var j: i64 = 0 88 while j < dv { 89 var acc: i64 = 0 90 var m: i64 = 0 91 while m < dk { acc = acc + kg_mul(M[i * dk + m], Sd[m * dv + j]); m = m + 1 } 92 S[i * dv + j] = acc + kg_mul(kg_mul(beta, k[i]), v[j]) 93 j = j + 1 94 } 95 i = i + 1 96 } 97 return 0 98} 99 100func main() -> i64 { 101 let ctr: *i64 = gv_ctr() 102 gv_head("nx_kda_gate -- the factorisation identity behind the KDA kernel, and the KDA:MLA schedule" as *u8) 103 104 let dk: i64 = 2 105 let dv: i64 = 2 106 let alpha: *i64 = sys_mmap(16) as *i64 107 let k: *i64 = sys_mmap(16) as *i64 108 let v: *i64 = sys_mmap(16) as *i64 109 let sc: *i64 = sys_mmap(16) as *i64 110 alpha[0] = KG_FP 111 alpha[1] = KG_FP 112 k[0] = KG_FP / 2 113 k[1] = KG_FP / 2 114 v[0] = KG_FP 115 v[1] = KG_FP * 2 116 117 // ---- T1 THE IDENTITY: factored O(dk*dv) == literal O(dk^2*dv), two steps, second from non-zero state 118 let Sa: *i64 = kg_state(dk, dv) 119 let Sb: *i64 = kg_state(dk, dv) 120 kg_step_factored(Sa, dk, dv, alpha, KG_FP, k, v, sc) 121 kg_step_naive(Sb, dk, dv, alpha, KG_FP, k, v) 122 kg_step_factored(Sa, dk, dv, alpha, KG_FP, k, v, sc) 123 kg_step_naive(Sb, dk, dv, alpha, KG_FP, k, v) 124 var t1: i64 = 1 125 var z: i64 = 0 126 while z < dk * dv { if Sa[z] != Sb[z] { t1 = 0 } z = z + 1 } 127 gv_check("T1 IDENTITY: the factored delta-rule update equals the literal (I - beta k k^T) form" as *u8, t1, ctr) 128 129 // ---- T2 NEG-CONTROL: T1 must have compared real values, not two piles of zeros 130 var t2: i64 = 0 131 if Sa[0] != 0 { if Sa[1] != Sa[0] { t2 = 1 } } 132 gv_check("T2 NEG-CONTROL: the compared state is non-zero and non-uniform" as *u8, t2, ctr) 133 134 // ---- T3 NON-DEGENERACY: with a basis-vector k the matrix would be diagonal and T1 would be trivial. 135 var t3: i64 = 0 136 let offdiag: i64 = 0 - kg_mul(kg_mul(KG_FP, k[0]), k[1]) 137 if offdiag != 0 { t3 = 1 } 138 gv_check("T3 NON-DEGENERACY: (I - beta k k^T) is genuinely non-diagonal, so T1 was not trivial" as *u8, t3, ctr) 139 140 // ---- T4 the 3:1 KDA:MLA interleave of Kimi Linear / K3 141 var t4: i64 = 1 142 if kdasched_kind(0, 3) != KDASCHED_LINEAR { t4 = 0 } 143 if kdasched_kind(1, 3) != KDASCHED_LINEAR { t4 = 0 } 144 if kdasched_kind(2, 3) != KDASCHED_LINEAR { t4 = 0 } 145 if kdasched_kind(3, 3) != KDASCHED_FULL { t4 = 0 } 146 if kdasched_kind(4, 3) != KDASCHED_LINEAR { t4 = 0 } 147 if kdasched_kind(7, 3) != KDASCHED_FULL { t4 = 0 } 148 gv_check("T4 3:1 interleave: layers 0,1,2 linear and layer 3 full, repeating" as *u8, t4, ctr) 149 150 // ---- T5 how many layers still grow a KV cache 151 var t5: i64 = 1 152 if kdasched_full_layers(64, 3) != 16 { t5 = 0 } 153 if kdasched_full_layers(64, 1) != 32 { t5 = 0 } 154 gv_check("T5 at 3:1 only 16 of 64 layers keep a growing KV cache" as *u8, t5, ctr) 155 156 // ---- T6 the reported 75% KV reduction, as a computed number rather than a comment 157 var t6: i64 = 1 158 if kdasched_kv_cut_permille(64, 3) != 750 { t6 = 0 } 159 if kdasched_kv_cut_permille(64, 1) != 500 { t6 = 0 } 160 gv_check("T6 the 3:1 schedule yields a 750-permille KV-growth cut (the reported 75%)" as *u8, t6, ctr) 161 162 // ---- T7 NEG-CONTROL: ratio=0 means no linear layers at all, so zero cut. Proves the schedule reads 163 // its ratio rather than returning a baked pattern. 164 var t7: i64 = 1 165 if kdasched_full_layers(64, 0) != 64 { t7 = 0 } 166 if kdasched_kv_cut_permille(64, 0) != 0 { t7 = 0 } 167 gv_check("T7 NEG-CONTROL: ratio=0 gives 64/64 full layers and a 0-permille cut" as *u8, t7, ctr) 168 169 let rc: i64 = gv_verdict("KDA-SCHED" as *u8, ctr, "factorisation identity proven non-degenerately; 3:1 schedule yields the 750-permille KV cut" as *u8) 170 return rc 171}