code wiki / (root) / nx_winograd_conv_test.nx

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}