nx_moe_gate.nx source
↩ module page · 174 lines · 6674 B
1// nx_moe_gate.nx -- MEASURED gate for the sovereign MoE sparse FFN
2// (nx_moe). PURE (no model, fast):
3//
4// NEG alloc rejects k > n_experts and zero dims
5// ROUTE organ's top-k == an INDEPENDENT largest-logits scan (same
6// tie rule); renormalized probs sum to 1 within 2^-15
7// ONE==FFN n_experts=1, top_k=1: nx_moe_forward BIT-EXACT equals the
8// bare expert FFN (softmax(single)=1.0 exactly, p*y=y)
9// SPARSE n=8, k=2, m=3: touch counters total EXACTLY m*k=6 -- the
10// sparsity is measured, not asserted
11// MIX forward's accumulation BIT-EXACT vs an oracle assembling
12// route + expert_ffn + p-weighted sum in the same documented
13// order
14// DET forward twice -> identical bits
15//
16// license_tier: ORIGINAL expect_exit: 0
17
18import "nx_syscalls.nx"
19import "nx_tier.nx"
20import "nx_f32.nx"
21import "nx_f32_cvt.nx"
22import "nx_f32_div.nx"
23import "nx_f32_softmax.nx"
24import "nx_f32_matmul_t.nx"
25import "nx_f32_activations.nx"
26import "nx_moe.nx"
27import "nx_fmt.nx"
28
29func mg_nl() -> i64 { fmt_puts("\n" as *u8); return 0 }
30func mg_lcg(s: i64) -> i64 { var v: i64 = s * 1103515245 + 12345; v = v & 2147483647; return v }
31func mg_fill(p: *i64, n: nx_int, state: *i64) -> i64 {
32 var s: i64 = state[0]
33 var i: nx_int = 0
34 while i < n {
35 s = mg_lcg(s)
36 let vv: i64 = nx_i32_to_f32((s % 9) - 4)
37 p[i] = vv
38 i = i + 1
39 }
40 state[0] = s
41 return 0
42}
43func mg_fill_layer(L: *NxMoeLayer, state: *i64) -> i64 {
44 mg_fill(L.W_router, L.n_experts * L.hidden_dim, state)
45 var e: nx_int = 0
46 while e < L.n_experts {
47 mg_fill(L.experts_gate[e] as *i64, L.ffn_dim * L.hidden_dim, state)
48 mg_fill(L.experts_up[e] as *i64, L.ffn_dim * L.hidden_dim, state)
49 mg_fill(L.experts_down[e] as *i64, L.hidden_dim * L.ffn_dim, state)
50 e = e + 1
51 }
52 return 0
53}
54func mg_cmp(a: *i64, b: *i64, n: nx_int) -> i64 {
55 var i: nx_int = 0
56 while i < n { if a[i] != b[i] { return 0 } i = i + 1 }
57 return 1
58}
59
60const MG_HD: nx_int = 8
61const MG_FD: nx_int = 16
62
63func main() -> i64 {
64 // ---- NEG: alloc guards -----------------------------------------
65 let bad1: *NxMoeLayer = nx_moe_layer_alloc(4, 5, MG_HD, MG_FD)
66 if bad1 != (0 as *NxMoeLayer) { return 11 }
67 let bad2: *NxMoeLayer = nx_moe_layer_alloc(0, 1, MG_HD, MG_FD)
68 if bad2 != (0 as *NxMoeLayer) { return 11 }
69 fmt_puts("MOE NEG alloc guards OK"); mg_nl()
70
71 let st: *i64 = sys_mmap(8) as *i64
72 st[0] = 20260709
73
74 // ---- ROUTE: independent top-k + renorm sum ----------------------
75 let L8: *NxMoeLayer = nx_moe_layer_alloc(8, 2, MG_HD, MG_FD)
76 mg_fill_layer(L8, st)
77 let x: *i64 = sys_mmap(MG_HD * 8) as *i64
78 mg_fill(x, MG_HD, st)
79 let idx: *i64 = sys_mmap(2 * 8) as *i64
80 let prb: *i64 = sys_mmap(2 * 8) as *i64
81 let vr: nx_int = nx_moe_route(L8, x, idx, prb)
82 if vr != NX_MOE_OK { return 21 }
83 // independent scan: top-2 of the raw logits (probs monotone in logits;
84 // same tie rule: strictly-greater, lowest index first).
85 let lg: *i64 = sys_mmap(8 * 8) as *i64
86 nx_f32_matmul_t(x, L8.W_router, lg, 1, MG_HD, 8)
87 var b1: nx_int = 0
88 var e1: nx_int = 1
89 while e1 < 8 { if nx_f32_gt(lg[e1], lg[b1]) == 1 { b1 = e1 } e1 = e1 + 1 }
90 var b2: nx_int = 0 - 1
91 var e2: nx_int = 0
92 while e2 < 8 {
93 if e2 != b1 {
94 var take: nx_int = 0
95 if b2 < 0 { take = 1 } else {
96 if nx_f32_gt(lg[e2], lg[b2]) == 1 { take = 1 }
97 }
98 if take == 1 { b2 = e2 }
99 }
100 e2 = e2 + 1
101 }
102 if (idx[0] as nx_int) != b1 { return 22 }
103 if (idx[1] as nx_int) != b2 { return 23 }
104 // renormalized probs sum to 1 within 2^-15.
105 let psum: i64 = nx_f32_add(prb[0], prb[1])
106 let diff: i64 = nx_f32_sub(psum, 0x3F800000) & 0x7FFFFFFF
107 if diff > 0x38000000 { return 24 }
108 fmt_puts("MOE ROUTE top-k==independent-scan + renorm-sum OK"); mg_nl()
109
110 // ---- ONE==FFN bit-exact ------------------------------------------
111 let L1: *NxMoeLayer = nx_moe_layer_alloc(1, 1, MG_HD, MG_FD)
112 mg_fill_layer(L1, st)
113 let y_ref: *i64 = sys_mmap(MG_HD * 8) as *i64
114 nx_moe_expert_ffn(L1, 0, x, y_ref)
115 let y_moe: *i64 = sys_mmap(MG_HD * 8) as *i64
116 let vf1: nx_int = nx_moe_forward(L1, x, 1, y_moe)
117 if vf1 != NX_MOE_OK { return 31 }
118 if mg_cmp(y_ref, y_moe, MG_HD) != 1 { return 32 }
119 fmt_puts("MOE ONE-expert == bare FFN BIT-EXACT OK"); mg_nl()
120
121 // ---- SPARSE: measured touch == m*k --------------------------------
122 let Ls: *NxMoeLayer = nx_moe_layer_alloc(8, 2, MG_HD, MG_FD)
123 mg_fill_layer(Ls, st)
124 let X3: *i64 = sys_mmap(3 * MG_HD * 8) as *i64
125 mg_fill(X3, 3 * MG_HD, st)
126 let O3: *i64 = sys_mmap(3 * MG_HD * 8) as *i64
127 let vf2: nx_int = nx_moe_forward(Ls, X3, 3, O3)
128 if vf2 != NX_MOE_OK { return 41 }
129 var tsum: i64 = 0
130 var te: nx_int = 0
131 while te < 8 {
132 tsum = tsum + Ls.touch[te]
133 if Ls.touch[te] > 3 { return 42 }
134 te = te + 1
135 }
136 if tsum != 6 { return 43 } // exactly m*k expert runs
137 fmt_puts("MOE SPARSE touch-total==m*k=6 (only top-k experts ran) OK"); mg_nl()
138
139 // ---- MIX: forward == oracle assembly ------------------------------
140 let x1: *i64 = ((X3 as i64) + 1 * MG_HD * 8) as *i64 // token 1
141 let oi: *i64 = sys_mmap(2 * 8) as *i64
142 let op: *i64 = sys_mmap(2 * 8) as *i64
143 nx_moe_route(Ls, x1, oi, op)
144 let yk: *i64 = sys_mmap(MG_HD * 8) as *i64
145 let expd: *i64 = sys_mmap(MG_HD * 8) as *i64
146 var d0: nx_int = 0
147 while d0 < MG_HD { expd[d0] = 0; d0 = d0 + 1 }
148 var kk: nx_int = 0
149 while kk < 2 {
150 nx_moe_expert_ffn(Ls, oi[kk] as nx_int, x1, yk)
151 var d: nx_int = 0
152 while d < MG_HD {
153 let contrib: i64 = nx_f32_mul(op[kk], yk[d])
154 let acc: i64 = nx_f32_add(expd[d], contrib)
155 expd[d] = acc
156 d = d + 1
157 }
158 kk = kk + 1
159 }
160 let o1: *i64 = ((O3 as i64) + 1 * MG_HD * 8) as *i64
161 if mg_cmp(expd, o1, MG_HD) != 1 { return 51 }
162 fmt_puts("MOE MIX forward == oracle assembly BIT-EXACT OK"); mg_nl()
163
164 // ---- DET: forward twice identical ---------------------------------
165 let O3b: *i64 = sys_mmap(3 * MG_HD * 8) as *i64
166 let vf3: nx_int = nx_moe_forward(Ls, X3, 3, O3b)
167 if vf3 != NX_MOE_OK { return 61 }
168 if mg_cmp(O3, O3b, 3 * MG_HD) != 1 { return 62 }
169 fmt_puts("MOE DET repeat-forward identical OK"); mg_nl()
170
171 fmt_puts("LIAR-KILL neg=1 route-independent=1 one==ffn=1 sparse-measured=1 mix-oracle=1 det=1"); mg_nl()
172 fmt_puts("MOE_GATE DONE"); mg_nl()
173 return 0
174}