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}