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}