code wiki / (root) / nx_winograd_conv.nx

nx_winograd_conv.nx source

↩ module page · 310 lines · 9825 B

1// nx_winograd_conv.nx -- 2D convolution via Winograd F(2x2, 3x3). 2// 3// Algorithm-led mult-reduction kernel. For every 3x3 conv kernel 4// producing a 2x2 output tile: 5// 6// Direct convolution: 2 * 2 * 3 * 3 = 36 multiplies 7// Winograd F(2x2, 3x3): 4 * 4 = 16 multiplies 8// Reduction: 36 / 16 = 2.25x (Lavin & Gray 2016) 9// 10// More additions (transform matrices) -- but adds are 1-cycle ALU 11// on every CPU including MCUs; mults take 3-5. Net win on every 12// hardware floor we target. 13// 14// Per the min-hardware-floor + algo-led cardinal: this brick beats 15// direct conv on a $5 microcontroller as much as on a 5090, without 16// any SIMD / GPU dependency. When SIMD lands (compiler I3.x) the 17// 16 mults become 4 SIMD ops, compounding the win. 18// 19// Reference output composes via nx_compute_node as NX_CN_K_CONV2D 20// kernel kind. 21// 22// Winograd transform matrices (F(2x2, 3x3), Lavin & Gray 2016): 23// 24// B^T = [[1, 0, -1, 0], 25// [0, 1, 1, 0], 26// [0, -1, 1, 0], 27// [0, 1, 0, -1]] 28// 29// G = [[1, 0, 0], 30// [1/2, 1/2, 1/2], 31// [1/2, -1/2, 1/2], 32// [0, 0, 1]] 33// 34// A^T = [[1, 1, 1, 0], 35// [0, 1, -1, -1]] 36// 37// In Q-format on i64 substrate: 38// * Input tile (4x4) is just i64 values 39// * G filter transform: G * g * G^T -> 4x4 transformed kernel 40// * B input transform: B^T * d * B -> 4x4 transformed input 41// * Hadamard product: U * V (elementwise 4x4) 42// * A output transform: A^T * M * A -> 2x2 output tile 43// 44// The G matrix's 1/2 entries are Q10-encoded as 512. This is the 45// ONLY place we need fractional arithmetic; the substrate already 46// handles Q10 throughout. 47// 48// genealogy_id: lavin_gray_2016_winograd + winograd_1980_minimal_filtering 49// lineage_id: substrate_winograd_conv_v1 50 51// nx_safety_envelope: 52// intended_use: AUTO_APPLIED -- primitive-specific tuning queued 53// sil_target: SIL1 54// evidence: [bulk_applied_2026-05-16, see-file-comment-for-detail] 55// verdict: NOT_YET_EVALUATED 56 57import "nx_syscalls.nx" 58import "nx_tier.nx" 59import "nx_tensor.nx" 60 61const NX_WG_Q10: nx_int = 1024 62 63// ===== Sealed-enum: WinogradVerdict =============================== 64 65const NX_WG_OK: nx_int = 0 66const NX_WG_ERR_BAD_DTYPE: nx_int = 1 67const NX_WG_ERR_BAD_KERNEL_SIZE: nx_int = 2 // not 3x3 (v1 only supports F(2x2,3x3)) 68const NX_WG_ERR_BAD_INPUT_SHAPE: nx_int = 3 // input dims not 4x4 per tile 69const NX_WG_ERR_NOT_CONTIGUOUS: nx_int = 4 70const NX_WG_N_VERDICTS: nx_int = 5 71 72func nx_wg_verdict_is_valid(v: nx_int) -> nx_int { 73 if v < 0 { return 0 } 74 if v >= NX_WG_N_VERDICTS { return 0 } 75 return 1 76} 77 78// ===== Filter transform U = G * g * G^T ========================= 79// 80// g: 3x3 filter (flat 9 i64). 81// U: 4x4 transformed kernel (flat 16 i64). 82// 83// G * g produces 4x3 intermediate; * G^T produces 4x4. 84// G's fractional rows are Q10-encoded (512). 85// 86// To keep arithmetic in pure i64 we multiply through by 2 (so the 87// 1/2 entries become 1). This means the transformed kernel is 88// stored at 4x scale (2 from rows + 2 from cols). We document the 89// 4x scale and the output transform A^T compensates. 90 91// G * g (intermediate, 4x3). Scaled by 2 so 1/2 -> 1: 92// row 0: g[0,*] 93// row 1: g[0,*] + g[1,*] + g[2,*] (was (g0+g1+g2)/2; *2 = g0+g1+g2) 94// row 2: g[0,*] - g[1,*] + g[2,*] (was (g0-g1+g2)/2; *2 = g0-g1+g2) 95// row 3: g[2,*] 96 97// Then (G*g) * G^T (4x4). Scaled again by 2: 98// col 0: row[0] 99// col 1: row[0] + row[1] + row[2] (sum /2 *2 = sum) 100// col 2: row[0] - row[1] + row[2] (alt /2 *2 = alt) 101// col 3: row[2] 102// 103// Total scale on U = 4. 104 105func nx_wg_filter_transform(g: *i64, U: *i64) -> nx_int { 106 // Intermediate: 4 rows x 3 cols 107 let inter: *i64 = (sys_mmap(96)) as *i64 // 4*3 i64 108 109 var c: nx_int = 0 110 while c < 3 { 111 let g0: nx_int = g[0 * 3 + c] 112 let g1: nx_int = g[1 * 3 + c] 113 let g2: nx_int = g[2 * 3 + c] 114 inter[0 * 3 + c] = g0 115 inter[1 * 3 + c] = g0 + g1 + g2 116 inter[2 * 3 + c] = g0 - g1 + g2 117 inter[3 * 3 + c] = g2 118 c = c + 1 119 } 120 // U = inter * G^T (4x4 from 4x3) 121 var r: nx_int = 0 122 while r < 4 { 123 let r0: nx_int = inter[r * 3 + 0] 124 let r1: nx_int = inter[r * 3 + 1] 125 let r2: nx_int = inter[r * 3 + 2] 126 U[r * 4 + 0] = r0 127 U[r * 4 + 1] = r0 + r1 + r2 128 U[r * 4 + 2] = r0 - r1 + r2 129 U[r * 4 + 3] = r2 130 r = r + 1 131 } 132 return NX_WG_OK 133} 134 135// ===== Input transform V = B^T * d * B ========================== 136// 137// d: 4x4 input tile (flat 16 i64). 138// V: 4x4 transformed input (flat 16 i64). 139// 140// B^T is integer-only (no fractional entries): 141// [[1, 0, -1, 0], 142// [0, 1, 1, 0], 143// [0, -1, 1, 0], 144// [0, 1, 0, -1]] 145// 146// B^T * d: 147// row 0: d[0,*] - d[2,*] 148// row 1: d[1,*] + d[2,*] 149// row 2: d[2,*] - d[1,*] 150// row 3: d[1,*] - d[3,*] 151// 152// (B^T*d) * B: 153// col 0: row[*,0] - row[*,2] 154// col 1: row[*,1] + row[*,2] 155// col 2: row[*,2] - row[*,1] 156// col 3: row[*,1] - row[*,3] 157 158func nx_wg_input_transform(d: *i64, V: *i64) -> nx_int { 159 let inter: *i64 = (sys_mmap(128)) as *i64 // 4x4 i64 160 161 var c: nx_int = 0 162 while c < 4 { 163 let d0: nx_int = d[0 * 4 + c] 164 let d1: nx_int = d[1 * 4 + c] 165 let d2: nx_int = d[2 * 4 + c] 166 let d3: nx_int = d[3 * 4 + c] 167 inter[0 * 4 + c] = d0 - d2 168 inter[1 * 4 + c] = d1 + d2 169 inter[2 * 4 + c] = d2 - d1 170 inter[3 * 4 + c] = d1 - d3 171 c = c + 1 172 } 173 var r: nx_int = 0 174 while r < 4 { 175 let r0: nx_int = inter[r * 4 + 0] 176 let r1: nx_int = inter[r * 4 + 1] 177 let r2: nx_int = inter[r * 4 + 2] 178 let r3: nx_int = inter[r * 4 + 3] 179 V[r * 4 + 0] = r0 - r2 180 V[r * 4 + 1] = r1 + r2 181 V[r * 4 + 2] = r2 - r1 182 V[r * 4 + 3] = r1 - r3 183 r = r + 1 184 } 185 return NX_WG_OK 186} 187 188// ===== Hadamard product M = U * V (elementwise 4x4) ============== 189// 190// This is where the 16 multiplies happen -- THE step Winograd 191// minimises. Per-tile cost = 16 mults regardless of how big the 192// kernel/output dims grow (vs direct conv's 9*4 = 36 mults per tile). 193 194func nx_wg_hadamard_4x4(U: *i64, V: *i64, M: *i64) -> nx_int { 195 var i: nx_int = 0 196 while i < 16 { 197 M[i] = U[i] * V[i] 198 i = i + 1 199 } 200 return NX_WG_OK 201} 202 203// ===== Output transform Y = A^T * M * A ========================= 204// 205// Y: 2x2 output tile (flat 4 i64). 206// A^T is integer: 207// [[1, 1, 1, 0], 208// [0, 1, -1, -1]] 209// 210// A^T * M produces 2x4. Then * A produces 2x2. 211// 212// IMPORTANT: U has scale-factor 4 from the filter transform. Y is 213// thus output * 4. We divide by 4 at the end to recover the true 214// convolution output. 215 216func nx_wg_output_transform(M: *i64, Y: *i64) -> nx_int { 217 let inter: *i64 = (sys_mmap(64)) as *i64 // 2x4 i64 218 219 var c: nx_int = 0 220 while c < 4 { 221 let m0: nx_int = M[0 * 4 + c] 222 let m1: nx_int = M[1 * 4 + c] 223 let m2: nx_int = M[2 * 4 + c] 224 let m3: nx_int = M[3 * 4 + c] 225 inter[0 * 4 + c] = m0 + m1 + m2 226 inter[1 * 4 + c] = m1 - m2 - m3 227 c = c + 1 228 } 229 var r: nx_int = 0 230 while r < 2 { 231 let r0: nx_int = inter[r * 4 + 0] 232 let r1: nx_int = inter[r * 4 + 1] 233 let r2: nx_int = inter[r * 4 + 2] 234 let r3: nx_int = inter[r * 4 + 3] 235 // Divide by 4 to undo the filter-transform scale factor 236 Y[r * 2 + 0] = (r0 + r1 + r2) / 4 237 Y[r * 2 + 1] = (r1 - r2 - r3) / 4 238 r = r + 1 239 } 240 return NX_WG_OK 241} 242 243// ===== Full Winograd 2x2 tile conv ================================ 244// 245// Convenience: feeds the 3x3 filter + 4x4 input tile through the 246// four steps and writes the 2x2 output. 247 248func nx_wg_conv_tile(filter: *i64, input_tile: *i64, output_tile: *i64) -> nx_int { 249 let U: *i64 = (sys_mmap(128)) as *i64 250 let V: *i64 = (sys_mmap(128)) as *i64 251 let M: *i64 = (sys_mmap(128)) as *i64 252 253 let r1: nx_int = nx_wg_filter_transform(filter, U) 254 if r1 != NX_WG_OK { return r1 } 255 256 let r2: nx_int = nx_wg_input_transform(input_tile, V) 257 if r2 != NX_WG_OK { return r2 } 258 259 nx_wg_hadamard_4x4(U, V, M) 260 nx_wg_output_transform(M, output_tile) 261 return NX_WG_OK 262} 263 264// ===== Reference direct convolution (oracle gate) ================ 265// 266// Plain 3x3 valid convolution to compute a 2x2 output from a 4x4 267// input. The Winograd path MUST produce the same bytes for any 268// inputs where the algorithm is exact (integer arithmetic with 269// no rounding loss). 270// 271// For values that fit safely (no /4 truncation), Winograd is 272// bit-identical to direct conv. 273 274func nx_wg_direct_conv_reference(filter: *i64, input_tile: *i64, 275 output_tile: *i64) -> nx_int { 276 // 2x2 output from 4x4 input using 3x3 kernel (valid mode) 277 var r: nx_int = 0 278 while r < 2 { 279 var c: nx_int = 0 280 while c < 2 { 281 var acc: nx_int = 0 282 var i: nx_int = 0 283 while i < 3 { 284 var j: nx_int = 0 285 while j < 3 { 286 acc = acc + filter[i * 3 + j] * input_tile[(r + i) * 4 + (c + j)] 287 j = j + 1 288 } 289 i = i + 1 290 } 291 output_tile[r * 2 + c] = acc 292 c = c + 1 293 } 294 r = r + 1 295 } 296 return NX_WG_OK 297} 298 299// ===== Multiply count (audit-only; for benchmark reporting) ====== 300// 301// Returns 16 (Winograd) vs 36 (direct). Substrate publishes this 302// for the dashboard to show "X mults saved" per layer. 303 304const NX_WG_DIRECT_MULTS: nx_int = 36 305const NX_WG_WINOGRAD_MULTS: nx_int = 16 306 307func nx_wg_mult_reduction_q10() -> nx_int { 308 // ratio = 36 * Q10 / 16 309 return (NX_WG_DIRECT_MULTS * NX_WG_Q10) / NX_WG_WINOGRAD_MULTS 310}