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}