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}