code wiki / (root) / nx_conv_gemm_gate.nx

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}