code wiki / _hdl_build / nx_f32_hw_matmul_gate.nx

nx_f32_hw_matmul_gate.nx source

↩ module page · 94 lines · 3817 B

1// nx_f32_hw_matmul_gate.nx -- proves the hardware-float GEMM is correct. 2// 3// T1 kat : hw 2x2 [[1,2],[3,4]]@[[5,6],[7,8]] == [[19,22],[43,50]] 4// T2 differential : hw == software nx_f32_matmul BIT-FOR-BIT on an 8x8 (binary32 cross-check) 5// T3 bad-dim : m=0 -> NX_HWMM_ERR 6// 7// GREEN only if all three hold. Evidence -> knowledge/status/f32_hw_matmul_gate.log. 8// Sovereign: imports nx_f32_hw_matmul (hw, -> nx_f32_hw SSE) + nx_f32_matmul (sw ref, -> nx_f32) 9// + nx_framed_append + nx_syscalls. license_tier: ORIGINAL 10import "nx_f32_hw_matmul.nx" 11import "nx_f32_matmul.nx" 12import "nx_framed_append.nx" 13import "nx_syscalls.nx" 14 15const HG_LOG: *u8 = "knowledge/status/f32_hw_matmul_gate.log" 16 17func hg_cat(dst: *u8, off: i64, s: *u8) -> i64 { var i: i64 = 0; while s[i] != 0 as u8 { dst[off + i] = s[i]; i = i + 1 } return off + i } 18func hg_catn(dst: *u8, off: i64, v: i64) -> i64 { 19 var m: i64 = v; var o: i64 = off 20 if m < 0 { m = 0 - m } 21 let t: *u8 = sys_mmap(28); var k: i64 = 0 22 if m == 0 { t[0] = 48 as u8; k = 1 } 23 while m > 0 { t[k] = (48 + (m % 10)) as u8; m = m / 10; k = k + 1 } 24 var i: i64 = 0 25 while i < k { dst[o + i] = t[k - 1 - i]; i = i + 1 } 26 return o + k 27} 28func hg_row(name: *u8, pass: i64) -> i64 { 29 let buf: *u8 = sys_mmap(528) 30 var o: i64 = 0 31 o = hg_cat(buf, o, "HWMM row=\x00" as *u8) 32 o = hg_cat(buf, o, name) 33 if pass == 1 { o = hg_cat(buf, o, " verdict=PASS\x00" as *u8) } else { o = hg_cat(buf, o, " verdict=FAIL\x00" as *u8) } 34 buf[o] = 0 as u8 35 fa_appendz(HG_LOG, buf, 512) 36 hmm_puts(" "); hmm_puts(name) 37 if pass == 1 { hmm_puts(" PASS\n") } else { hmm_puts(" FAIL\n") } 38 return 0 39} 40 41func main() -> i64 { 42 hmm_puts("f32-hw-matmul gate (hardware SSE GEMM: KAT + differential-vs-software)\n") 43 44 // T1 KAT 45 let A: *i64 = sys_mmap(4 * 8) as *i64 46 A[0] = f32_of(1); A[1] = f32_of(2); A[2] = f32_of(3); A[3] = f32_of(4) 47 let B: *i64 = sys_mmap(4 * 8) as *i64 48 B[0] = f32_of(5); B[1] = f32_of(6); B[2] = f32_of(7); B[3] = f32_of(8) 49 let C: *i64 = sys_mmap(4 * 8) as *i64 50 nx_f32_hw_matmul(A, B, C, 2, 2, 2) 51 var t1: i64 = 0 52 if f32_int(C[0]) == 19 { if f32_int(C[1]) == 22 { if f32_int(C[2]) == 43 { if f32_int(C[3]) == 50 { t1 = 1 } } } } 53 54 // T2 differential: 8x8 small-int inputs -> hw == sw bit-for-bit 55 let N: i64 = 8 56 let DA: *i64 = sys_mmap(N * N * 8) as *i64 57 let DB: *i64 = sys_mmap(N * N * 8) as *i64 58 let CHW: *i64 = sys_mmap(N * N * 8) as *i64 59 let CSW: *i64 = sys_mmap(N * N * 8) as *i64 60 var i: i64 = 0 61 while i < N * N { DA[i] = f32_of((i % 5) + 1); DB[i] = f32_of((i % 7) + 1); i = i + 1 } 62 nx_f32_hw_matmul(DA, DB, CHW, N, N, N) 63 nx_f32_matmul(DA, DB, CSW, N, N, N) 64 var same: i64 = 1; i = 0 65 while i < N * N { if CHW[i] != CSW[i] { same = 0; i = N * N } else { i = i + 1 } } 66 var t2: i64 = same 67 68 // T3 bad-dim 69 var t3: i64 = 0 70 if nx_f32_hw_matmul(A, B, C, 0, 2, 2) == NX_HWMM_ERR { t3 = 1 } 71 72 var passes: i64 = 0 73 if t1 == 1 { passes = passes + 1 } 74 if t2 == 1 { passes = passes + 1 } 75 if t3 == 1 { passes = passes + 1 } 76 var green: i64 = 0 77 if passes == 3 { green = 1 } 78 79 hg_row("T1-hw-kat-correct \x00" as *u8, t1) 80 hg_row("T2-differential-hw==sw \x00" as *u8, t2) 81 hg_row("T3-bad-dim-error \x00" as *u8, t3) 82 83 let vb: *u8 = sys_mmap(528) 84 var o: i64 = 0 85 o = hg_cat(vb, o, "F32-HW-MATMUL verdict=\x00" as *u8) 86 if green == 1 { o = hg_cat(vb, o, "GREEN\x00" as *u8) } else { o = hg_cat(vb, o, "RED\x00" as *u8) } 87 o = hg_cat(vb, o, " passes=\x00" as *u8); o = hg_catn(vb, o, passes); o = hg_cat(vb, o, "/3 END\x00" as *u8) 88 vb[o] = 0 as u8 89 fa_appendz(HG_LOG, vb, 512) 90 hmm_puts(vb); hmm_puts("\n") 91 92 if green == 1 { return 0 } 93 return 1 94}