nx_winograd_conv_test.nx source
↩ module page · 136 lines · 4968 B
1// nx_winograd_conv_test.nx -- algo-led correctness: Winograd output
2// MUST match direct convolution bit-exact for integer inputs where
3// the algorithm is exact. Oracle gate.
4
5import "nx_syscalls.nx"
6import "nx_tier.nx"
7import "nx_winograd_conv.nx"
8import "nx_numeric_oracle.nx"
9import "nx_tensor.nx"
10
11func main() -> nx_int {
12 // ===== 3x3 filter ============================================
13 let filter: *i64 = (sys_mmap(72)) as *i64 // 9 i64
14 filter[0] = 1; filter[1] = 0; filter[2] = -1
15 filter[3] = 2; filter[4] = 0; filter[5] = -2
16 filter[6] = 1; filter[7] = 0; filter[8] = -1
17 // (Sobel-X style; pure integer, well-behaved)
18
19 // ===== 4x4 input tile ========================================
20 let input_tile: *i64 = (sys_mmap(128)) as *i64
21 var i: nx_int = 0
22 while i < 16 {
23 // values 0..15 multiplied by 4 (well within Winograd
24 // exactness range; avoids /4 truncation)
25 input_tile[i] = i * 4
26 i = i + 1
27 }
28
29 // ===== Direct convolution (reference) =======================
30 let direct_out: *i64 = (sys_mmap(32)) as *i64
31 nx_wg_direct_conv_reference(filter, input_tile, direct_out)
32
33 // Hand-verify one cell. Filter = Sobel-X applied at output[0,0]:
34 // Input window:
35 // [0, 4, 8]
36 // [16, 20, 24]
37 // [32, 36, 40]
38 // Sum = 1*0 + 0*4 + (-1)*8 + 2*16 + 0*20 + (-2)*24 + 1*32 + 0*36 + (-1)*40
39 // = 0 + 0 - 8 + 32 + 0 - 48 + 32 + 0 - 40 = -32
40 if direct_out[0] != 0 - 32 { return 1 }
41
42 // ===== Winograd convolution =================================
43 let winograd_out: *i64 = (sys_mmap(32)) as *i64
44 let rc: nx_int = nx_wg_conv_tile(filter, input_tile, winograd_out)
45 if rc != NX_WG_OK { return 2 }
46
47 // ===== Oracle: Winograd == direct bit-exact =================
48 let witness: *i64 = (sys_mmap(NX_NO_WITNESS_FIELDS * 8)) as *i64
49 let sh: *i64 = (sys_mmap(16)) as *i64
50 sh[0] = 2; sh[1] = 2
51 let err: *i64 = (sys_mmap(8)) as *i64
52 let t_d: *NxTensor = nx_t_alloc(NX_DT_I64, sh, 2, err)
53 let t_w: *NxTensor = nx_t_alloc(NX_DT_I64, sh, 2, err)
54 let pd: *i64 = t_d.storage as *i64
55 let pw: *i64 = t_w.storage as *i64
56 var k: nx_int = 0
57 while k < 4 {
58 pd[k] = direct_out[k]
59 pw[k] = winograd_out[k]
60 k = k + 1
61 }
62 let v: nx_int = nx_no_check_bit_exact_i64(t_d, t_w, witness)
63 if v != NX_NO_VERDICT_EQUAL {
64 // Witness has the failing cell -- the test reports 10+verdict
65 return 10 + v
66 }
67
68 // ===== Verify intermediate transforms ========================
69 //
70 // Identity filter [[0,0,0],[0,1,0],[0,0,0]]: Winograd output
71 // should reproduce the centre-cropped 2x2 of the input.
72 let id_filter: *i64 = (sys_mmap(72)) as *i64
73 var f: nx_int = 0
74 while f < 9 { id_filter[f] = 0; f = f + 1 }
75 id_filter[4] = 1
76 let id_out: *i64 = (sys_mmap(32)) as *i64
77 nx_wg_conv_tile(id_filter, input_tile, id_out)
78 // input[1,1] = 5*4 = 20
79 if id_out[0] != 20 { return 20 }
80 // input[1,2] = 6*4 = 24
81 if id_out[1] != 24 { return 21 }
82 // input[2,1] = 9*4 = 36
83 if id_out[2] != 36 { return 22 }
84 // input[2,2] = 10*4 = 40
85 if id_out[3] != 40 { return 23 }
86
87 // ===== Verify zero filter -> zero output ====================
88 let zero_filter: *i64 = (sys_mmap(72)) as *i64
89 var z: nx_int = 0
90 while z < 9 { zero_filter[z] = 0; z = z + 1 }
91 let zero_out: *i64 = (sys_mmap(32)) as *i64
92 nx_wg_conv_tile(zero_filter, input_tile, zero_out)
93 var zc: nx_int = 0
94 while zc < 4 {
95 if zero_out[zc] != 0 { return 30 + zc }
96 zc = zc + 1
97 }
98
99 // ===== Multiply-count reduction ratio ======================
100 //
101 // 36 / 16 in Q10 = 36864 / 16 = 2304 -> 2.25x
102 if nx_wg_mult_reduction_q10() != 2304 { return 40 }
103
104 // ===== Sealed verdict bands =================================
105 if nx_wg_verdict_is_valid(NX_WG_OK) != 1 { return 50 }
106 if nx_wg_verdict_is_valid(NX_WG_N_VERDICTS) != 0 { return 51 }
107 if nx_wg_verdict_is_valid(0 - 1) != 0 { return 52 }
108
109 // ===== Bigger / negative-filter case for variety ===========
110 //
111 // Filter = [[1, -1, 1], [-1, 2, -1], [1, -1, 1]] -- arbitrary
112 // mix; non-trivial; must still match direct conv.
113 let filter2: *i64 = (sys_mmap(72)) as *i64
114 filter2[0] = 1; filter2[1] = -1; filter2[2] = 1
115 filter2[3] = -1; filter2[4] = 2; filter2[5] = -1
116 filter2[6] = 1; filter2[7] = -1; filter2[8] = 1
117
118 let input2: *i64 = (sys_mmap(128)) as *i64
119 var ii: nx_int = 0
120 while ii < 16 {
121 input2[ii] = (ii * 8) - 30 // -30..90 range, scaled by 4 implicitly
122 ii = ii + 1
123 }
124
125 let d2: *i64 = (sys_mmap(32)) as *i64
126 let w2: *i64 = (sys_mmap(32)) as *i64
127 nx_wg_direct_conv_reference(filter2, input2, d2)
128 nx_wg_conv_tile(filter2, input2, w2)
129 var bb: nx_int = 0
130 while bb < 4 {
131 if d2[bb] != w2[bb] { return 60 + bb }
132 bb = bb + 1
133 }
134
135 return 0
136}