nx_conv_gemm_gate.nx source
↩ module page · 156 lines · 7928 B
1// nx_conv_gemm_gate.nx -- wire the COMPUTE LEVER into the diffusion workhorse: f32 Conv2D via im2col->GEMM.
2// Convolution dominates diffusion U-Nets (what sdcpp/GGML run). The fast path (what GGML uses) is im2col
3// then GEMM -- so a 3x3 conv becomes Weight[Cout, Cin*9] x im2col_T[H*W, Cin*9]^T, reusing the FMA-4acc
4// matmul kernel built this session. This compares a NAIVE direct f32 conv (scalar __f32_mul/add) against
5// im2col + FMA-4acc GEMM, BIT-EXACT (small-int f32) + MEASURES the speedup. This is the first sovereign
6// f32 image-gen primitive accelerated by our AVX2/FMA stack -- the concrete start of the sdcpp benchmark path.
7// No hw writes (Rule 26). expect_exit: 0 license_tier: ORIGINAL
8import "nx_syscalls.nx"
9import "nx_gate_verdict.nx"
10
11func cg_puts(s: *u8) -> i64 { var n: i64=0; while s[n]!=(0 as u8){n=n+1} sys_write(1,s,n); return 0 }
12func cg_num(v: i64) -> i64 { let b: *u8=sys_mmap(28); var m: i64=v; if m<0{m=0-m;sys_write(1,"-" as *u8,1)} let t: *u8=sys_mmap(28); var k: i64=0; if m==0{t[0]=48 as u8;k=1} while m>0{t[k]=(48+(m%10)) as u8;m=m/10;k=k+1} var i: i64=0; while i<k{b[i]=t[k-1-i];i=i+1} sys_write(1,b,k); return 0 }
13func pack4(buf: *u8, idx: i64, bits: i64) -> i64 { buf[idx*4+0]=(bits) as u8; buf[idx*4+1]=(bits>>8) as u8; buf[idx*4+2]=(bits>>16) as u8; buf[idx*4+3]=(bits>>24) as u8; return 0 }
14
15// 4-accumulator FMA dot over K packed f32 (K divisible by 32); wr/ir = byte addresses
16func fma4dot(wr: i64, ir: i64, K: i64, a0: *u8, a1: *u8, a2: *u8, a3: *u8) -> i64 {
17 let z0: *i64=a0 as *i64
18 let z1: *i64=a1 as *i64
19 let z2: *i64=a2 as *i64
20 let z3: *i64=a3 as *i64
21 z0[0]=0; z0[1]=0; z0[2]=0; z0[3]=0
22 z1[0]=0; z1[1]=0; z1[2]=0; z1[3]=0
23 z2[0]=0; z2[1]=0; z2[2]=0; z2[3]=0
24 z3[0]=0; z3[1]=0; z3[2]=0; z3[3]=0
25 var k: i64=0
26 while k<K {
27 __f32x8_fma(a0, (wr+k*4) as *u8, (ir+k*4) as *u8)
28 __f32x8_fma(a1, (wr+(k+8)*4) as *u8, (ir+(k+8)*4) as *u8)
29 __f32x8_fma(a2, (wr+(k+16)*4) as *u8, (ir+(k+16)*4) as *u8)
30 __f32x8_fma(a3, (wr+(k+24)*4) as *u8, (ir+(k+24)*4) as *u8)
31 k=k+32
32 }
33 let lo: i64=__f32_add(__f32x8_hsum(a0), __f32x8_hsum(a1))
34 let hi: i64=__f32_add(__f32x8_hsum(a2), __f32x8_hsum(a3))
35 return __f32_add(lo, hi)
36}
37
38func main() -> i64 {
39 cg_puts("f32 Conv2D via im2col -> FMA-GEMM (the diffusion workhorse, sdcpp's fast path) vs naive direct conv\n\n" as *u8)
40 let CIN: i64=32
41 let COUT: i64=32
42 let HH: i64=32
43 let WW: i64=32
44 let HW: i64=HH*WW
45 let KG: i64=CIN*9 // contraction dim = Cin * 3 * 3 = 288 (div by 32)
46
47 let inp: *i64 = sys_mmap(CIN*HH*WW*8) as *i64 // input, i64-f32
48 let wt: *i64 = sys_mmap(COUT*KG*8) as *i64 // weights [Cout][Cin*9], i64-f32
49 let wp: *u8 = sys_mmap(COUT*KG*4) // weights packed
50 let im: *u8 = sys_mmap(HW*KG*4) // im2col_T [H*W][Cin*9] packed
51 let outn: *i64 = sys_mmap(COUT*HW*8) as *i64 // naive output
52 let outg: *i64 = sys_mmap(COUT*HW*8) as *i64 // gemm output
53 let a0: *u8=sys_mmap(64)
54 let a1: *u8=sys_mmap(64)
55 let a2: *u8=sys_mmap(64)
56 let a3: *u8=sys_mmap(64)
57
58 // fill input + weights with small ints (exact in f32)
59 var c: i64=0
60 while c<CIN { var h: i64=0
61 while h<HH { var w: i64=0
62 while w<WW { inp[(c*HH+h)*WW+w]=__f32_from_i64(((c+h+w)%4)+1); w=w+1 }
63 h=h+1 }
64 c=c+1 }
65 var co: i64=0
66 while co<COUT { var kk: i64=0
67 while kk<KG { wt[co*KG+kk]=__f32_from_i64(((co+kk)%3)+1); pack4(wp, co*KG+kk, wt[co*KG+kk]); kk=kk+1 }
68 co=co+1 }
69
70 // im2col_T: row p=(oh,ow), col k=c*9+kh*3+kw -> padded input (pad 1, stride 1); out-of-bounds = 0
71 var oh: i64=0
72 while oh<HH { var ow: i64=0
73 while ow<WW {
74 let p: i64=oh*WW+ow
75 var c2: i64=0
76 while c2<CIN { var kh: i64=0
77 while kh<3 { var kw: i64=0
78 while kw<3 {
79 let ih: i64=oh+kh-1
80 let iw: i64=ow+kw-1
81 var v: i64=__f32_from_i64(0)
82 if ih>=0 { if ih<HH { if iw>=0 { if iw<WW { v=inp[(c2*HH+ih)*WW+iw] } } } }
83 pack4(im, p*KG + (c2*9+kh*3+kw), v)
84 kw=kw+1 }
85 kh=kh+1 }
86 c2=c2+1 }
87 ow=ow+1 }
88 oh=oh+1 }
89
90 // NAIVE direct conv (scalar f32) -- same math, the baseline
91 let t0: i64=sys_now_us()
92 co=0
93 while co<COUT { var oh2: i64=0
94 while oh2<HH { var ow2: i64=0
95 while ow2<WW {
96 var acc: i64=__f32_from_i64(0)
97 var c3: i64=0
98 while c3<CIN { var kh2: i64=0
99 while kh2<3 { var kw2: i64=0
100 while kw2<3 {
101 let ih2: i64=oh2+kh2-1
102 let iw2: i64=ow2+kw2-1
103 if ih2>=0 { if ih2<HH { if iw2>=0 { if iw2<WW {
104 acc=__f32_add(acc, __f32_mul(wt[co*KG+(c3*9+kh2*3+kw2)], inp[(c3*HH+ih2)*WW+iw2])) } } } }
105 kw2=kw2+1 }
106 kh2=kh2+1 }
107 c3=c3+1 }
108 outn[co*HW+oh2*WW+ow2]=acc
109 ow2=ow2+1 }
110 oh2=oh2+1 }
111 co=co+1 }
112 let t1: i64=sys_now_us()
113
114 // im2col + FMA-GEMM conv
115 let wpb: i64=wp as i64
116 let imb: i64=im as i64
117 co=0
118 while co<COUT { var p2: i64=0
119 while p2<HW { outg[co*HW+p2]=fma4dot(wpb+co*KG*4, imb+p2*KG*4, KG, a0, a1, a2, a3); p2=p2+1 }
120 co=co+1 }
121 let t2: i64=sys_now_us()
122
123 var mism: i64=0
124 var z: i64=0
125 while z<COUT*HW { if __f32_to_i64(outn[z]) != __f32_to_i64(outg[z]) { mism=mism+1 } z=z+1 }
126
127 var us_n: i64=t1-t0
128 if us_n<=0 { us_n=1 }
129 var us_g: i64=t2-t1
130 if us_g<=0 { us_g=1 }
131 let sp: i64=us_n*100/us_g
132 let macs: i64=COUT*HW*KG
133 let mf_g: i64=2*macs/(us_g+1)
134
135 cg_puts(" conv: "); cg_num(CIN); cg_puts("ch "); cg_num(HH); cg_puts("x"); cg_num(WW); cg_puts(" -> "); cg_num(COUT); cg_puts("ch, 3x3 pad1 (GEMM K="); cg_num(KG); cg_puts(", "); cg_num(COUT*HW); cg_puts(" outputs)\n");
136 cg_puts(" out[0]: naive="); cg_num(__f32_to_i64(outn[0])); cg_puts(" gemm="); cg_num(__f32_to_i64(outg[0])); cg_puts("\n");
137 cg_puts(" bit-exact mismatches (gemm vs naive): "); cg_num(mism); cg_puts(" / "); cg_num(COUT*HW); cg_puts("\n");
138 cg_puts(" naive direct = "); cg_num(us_n); cg_puts("us im2col+FMA-GEMM = "); cg_num(us_g); cg_puts("us SPEEDUP = "); cg_num(sp/100); cg_puts("."); cg_num((sp%100)/10); cg_num(sp%10); cg_puts("x ("); cg_num(mf_g); cg_puts(" MFLOP/s)\n\n");
139
140 var pass: i64=0
141 var ttl: i64=0
142 ttl=ttl+1; cg_puts(" T1 GEMM conv == naive direct conv BIT-EXACT (im2col + FMA correct): "); if mism==0 { pass=pass+1; cg_puts("PASS\n") } else { cg_puts("FAIL\n") }
143 ttl=ttl+1; cg_puts(" T2 the compute lever ACCELERATES conv (GEMM faster than naive): "); if us_g<us_n { pass=pass+1; cg_puts("PASS\n") } else { cg_puts("FAIL\n") }
144 ttl=ttl+1; cg_puts(" T3 sovereign f32 conv primitive runs (the sdcpp-path bridge): "); if mf_g>0 { pass=pass+1; cg_puts("PASS\n") } else { cg_puts("FAIL\n") }
145
146 cg_puts("NX-CONV-GEMM-GATE passed "); cg_num(pass); cg_puts("/"); cg_num(ttl)
147 // MIGRATED onto nx_gate_verdict by nx_gate_dry_apply (D001, minimal form): every check
148 // row above is untouched, so the PASS/FAIL vector cannot change; only the hand-rolled
149 // verdict emission is replaced by the ONE shared base class. Proven by nx_gate_migrate verify.
150 let ctr__dry: *i64 = gv_ctr()
151 ctr__dry[0] = pass
152 ctr__dry[1] = ttl
153 let rc__dry: i64 = gv_verdict("CONV-GEMM-GATE" as *u8, ctr__dry, "f32 conv via im2col+FMA-GEMM, bit-exact + accelerated -- the diffusion workhorse on our compute lever)" as *u8)
154 sys_exit(rc__dry)
155 return rc__dry
156}