code wiki / _hdl_build / _ce_grad_gate_authored.nx
_ce_grad_gate_authored.nx source
↩ module page · 106 lines · 4630 B
1// _ce_grad_gate_authored.nx -- the VERIFIED-GRADIENT gate for tg_celoss (op 6, fused
2// softmax+cross-entropy), per the law established at T7: every op lands WITH its gradcheck.
3// GATES: A gradcheck (analytic == central finite difference, all 5 logits, h=1/128, rel<1/32
4// floor 1/64) | B KAT uniform logits -> CE == ln(n) (the mathematical anchor, oracle =
5// nx_f32_log itself on a DIFFERENT path: log(5) vs logsumexp of equal shifts) | C conservation
6// (softmax - onehot sums to 0 => analytic grads sum to ~0). LAWS: flat ifs, no &&/||.
7// license_tier: ORIGINAL
8import "nx_tgrad_core.nx"
9
10// plain (tape-free) CE forward for the finite-difference oracle
11func ce_fwd(pv: *i64, tv: *i64, n: i64) -> i64 {
12 var mx: i64 = pv[0]
13 var i: i64 = 1
14 while i < n {
15 if nx_f32_gt(pv[i], mx) == 1 { mx = pv[i] }
16 i = i + 1
17 }
18 var s: i64 = 0
19 var dot: i64 = 0
20 i = 0
21 while i < n {
22 let sh: i64 = nx_f32_sub(pv[i], mx)
23 s = nx_f32_add(s, nx_f32_exp(sh))
24 dot = nx_f32_add(dot, nx_f32_mul(tv[i], sh))
25 i = i + 1
26 }
27 return nx_f32_sub(nx_f32_log(s), dot)
28}
29
30func main() -> i64 {
31 _tg_puts("=== CE GRADCHECK GATE (tg_celoss op 6: fused softmax+cross-entropy, verified-gradient law) ===\n" as *u8)
32 let tape: *i64 = sys_mmap(32768) as *i64
33 let nb: *i64 = sys_mmap(16) as *i64
34 let arena: *i64 = sys_mmap(131072) as *i64
35 let ab: *i64 = sys_mmap(16) as *i64
36 let pv: *i64 = sys_mmap(64) as *i64
37 let tv: *i64 = sys_mmap(64) as *i64
38 // varied logits (positive + negative, distinct), one-hot at class 2
39 pv[0] = tg_q(3, 4)
40 pv[1] = nx_f32_neg(tg_q(1, 2))
41 pv[2] = tg_q(5, 4)
42 pv[3] = tg_q(1, 8)
43 pv[4] = nx_f32_neg(tg_q(7, 8))
44 var i: i64 = 0
45 while i < 5 { tv[i] = 0; i = i + 1 }
46 tv[2] = nx_i32_to_f32(1)
47 nb[0] = 0
48 ab[0] = 0
49 let lp: i64 = tg_leaf(tape, nb, arena, ab, pv, 5, 1)
50 let lt: i64 = tg_leaf(tape, nb, arena, ab, tv, 5, 1)
51 let loss: i64 = tg_celoss(tape, nb, arena, ab, lp, lt)
52 tg_backward(tape, nb[0], loss)
53 let ga: *i64 = tg_gradp(tape, lp)
54 let h128: i64 = tg_q(1, 128)
55 var pa: i64 = 1
56 var gsum: i64 = 0
57 var k: i64 = 0
58 while k < 5 {
59 gsum = nx_f32_add(gsum, ga[k])
60 let save: i64 = pv[k]
61 pv[k] = nx_f32_add(save, h128)
62 let fp: i64 = ce_fwd(pv, tv, 5)
63 pv[k] = nx_f32_sub(save, h128)
64 let fm: i64 = ce_fwd(pv, tv, 5)
65 pv[k] = save
66 let fdif: i64 = nx_f32_div(nx_f32_sub(fp, fm), tg_q(1, 64))
67 var den: i64 = nx_f32_abs(fdif)
68 if nx_f32_lt(den, tg_q(1, 64)) == 1 { den = tg_q(1, 64) }
69 let rel: i64 = nx_f32_div(nx_f32_abs(nx_f32_sub(ga[k], fdif)), den)
70 _tg_puts(" gradcheck logit" as *u8); _tg_num(k)
71 _tg_puts(" analytic-milli=" as *u8); _tg_num(tg_milli(ga[k]))
72 _tg_puts(" fd-milli=" as *u8); _tg_num(tg_milli(fdif))
73 _tg_puts(" rel-milli=" as *u8); _tg_num(tg_milli(rel)); _tg_puts("\n" as *u8)
74 if nx_f32_lt(rel, tg_q(1, 32)) == 0 { pa = 0 }
75 k = k + 1
76 }
77 if pa == 1 { _tg_puts(" GATE A CE gradcheck (5 logits): PASS\n" as *u8) } else { _tg_puts(" GATE A CE gradcheck: FAIL\n" as *u8) }
78 // GATE B: uniform logits -> CE = ln(n); oracle = nx_f32_log(5) via a different code path
79 let pu: *i64 = sys_mmap(64) as *i64
80 k = 0
81 while k < 5 { pu[k] = tg_q(1, 4); k = k + 1 }
82 let ce_u: i64 = ce_fwd(pu, tv, 5)
83 let ln5: i64 = nx_f32_log(nx_i32_to_f32(5))
84 var pb: i64 = 1
85 if nx_f32_lt(nx_f32_abs(nx_f32_sub(ce_u, ln5)), tg_q(1, 1000)) == 0 { pb = 0 }
86 _tg_puts(" uniform CE-milli=" as *u8); _tg_num(tg_milli(ce_u))
87 _tg_puts(" ln5-milli=" as *u8); _tg_num(tg_milli(ln5)); _tg_puts("\n" as *u8)
88 if pb == 1 { _tg_puts(" GATE B uniform-logits CE == ln(5): PASS\n" as *u8) } else { _tg_puts(" GATE B: FAIL\n" as *u8) }
89 // GATE C: conservation -- analytic grads sum to ~0 (softmax sums to 1, one-hot sums to 1)
90 var pc: i64 = 1
91 if nx_f32_lt(nx_f32_abs(gsum), tg_q(1, 1000)) == 0 { pc = 0 }
92 _tg_puts(" gradsum-milli=" as *u8); _tg_num(tg_milli(gsum)); _tg_puts("\n" as *u8)
93 if pc == 1 { _tg_puts(" GATE C gradient conservation (sum == 0): PASS\n" as *u8) } else { _tg_puts(" GATE C: FAIL\n" as *u8) }
94 var gates: i64 = 0
95 if pa == 1 { gates = gates + 1 }
96 if pb == 1 { gates = gates + 1 }
97 if pc == 1 { gates = gates + 1 }
98 if gates == 3 {
99 _tg_puts(" CE-GRADCHECK GATE: PASS (cross-entropy joins the verified-gradient op set)\n" as *u8)
100 sys_exit(0)
101 return 0
102 }
103 _tg_puts(" CE-GRADCHECK GATE: FAIL\n" as *u8)
104 sys_exit(1)
105 return 1
106}