code wiki / (root) / nx_moe_gate.nx

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}