code wiki / (root) / nx_fma_matmul_gate.nx

nx_fma_matmul_gate.nx source

↩ module page · 132 lines · 6703 B

1// nx_fma_matmul_gate.nx -- PROVES the FMA vector-accumulate kernel beats per-chunk-hsum AVX2. 2// __f32x8_dot hsums every 8 elements (vextractf128 + SSE reduce each chunk). __f32x8_fma instead does a 3// FUSED 8-wide multiply-add into a memory accumulator (vfmadd231ps, encoder proven nxasm_vex_kat 7/7) and 4// hsums ONCE per dot via __f32x8_hsum. Fewer ops/MAC + the fused mul-add = the compute lever the roofline 5// diagnostic pointed to (we are compute-bound at this size). Three matmuls, same data: scalar vs AVX2-dot 6// vs AVX2-FMA; all BIT-EXACT (small-int f32) + a known answer + measures the FMA win over plain AVX2. 7// No hw writes (Rule 26). expect_exit: 0 license_tier: ORIGINAL 8import "nx_syscalls.nx" 9import "nx_gate_verdict.nx" 10 11func gx_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 gx_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 15func mm_scalar(sa: *i64, sb: *i64, c: *i64, M: i64, N: i64, K: i64) -> i64 { 16 var i: i64=0 17 while i<M { var j: i64=0 18 while j<N { var acc: i64=__f32_from_i64(0); var k: i64=0 19 while k<K { acc=__f32_add(acc, __f32_mul(sa[i*K+k], sb[j*K+k])); k=k+1 } 20 c[i*N+j]=acc; j=j+1 } 21 i=i+1 } 22 return 0 23} 24func mm_dot(pa: *u8, pb: *u8, c: *i64, M: i64, N: i64, K: i64) -> i64 { 25 let pab: i64=pa as i64 26 let pbb: i64=pb as i64 27 var i: i64=0 28 while i<M { var j: i64=0 29 while j<N { var acc: i64=__f32_from_i64(0); var k: i64=0 30 while k<K { acc=__f32_add(acc, __f32x8_dot((pab+(i*K+k)*4) as *u8, (pbb+(j*K+k)*4) as *u8)); k=k+8 } 31 c[i*N+j]=acc; j=j+1 } 32 i=i+1 } 33 return 0 34} 35func mm_fma(pa: *u8, pb: *u8, c: *i64, acc: *u8, M: i64, N: i64, K: i64) -> i64 { 36 let pab: i64=pa as i64 37 let pbb: i64=pb as i64 38 let az: *i64 = acc as *i64 39 var i: i64=0 40 while i<M { var j: i64=0 41 while j<N { 42 az[0]=0; az[1]=0; az[2]=0; az[3]=0 // zero the 8-f32 accumulator 43 var k: i64=0 44 while k<K { __f32x8_fma(acc, (pab+(i*K+k)*4) as *u8, (pbb+(j*K+k)*4) as *u8); k=k+8 } 45 c[i*N+j]=__f32x8_hsum(acc); j=j+1 } 46 i=i+1 } 47 return 0 48} 49 50func main() -> i64 { 51 gx_puts("FMA vector-accumulate matmul: __f32x8_fma (fused, hsum ONCE) vs __f32x8_dot (hsum per 8) vs scalar\n\n" as *u8) 52 let M: i64=64 53 let N: i64=64 54 let K: i64=64 55 let sa: *i64 = sys_mmap(M*K*8) as *i64 56 let sb: *i64 = sys_mmap(N*K*8) as *i64 57 let pa: *u8 = sys_mmap(M*K*4) 58 let pb: *u8 = sys_mmap(N*K*4) 59 let acc: *u8 = sys_mmap(64) 60 let c0: *i64 = sys_mmap(M*N*8) as *i64 61 let cd: *i64 = sys_mmap(M*N*8) as *i64 62 let cf: *i64 = sys_mmap(M*N*8) as *i64 63 64 var i: i64=0 65 while i<M { var k: i64=0 66 while k<K { let v: i64=__f32_from_i64(((i+k)%4)+1); sa[i*K+k]=v; pack4(pa,i*K+k,v); k=k+1 } 67 i=i+1 } 68 var j: i64=0 69 while j<N { var k2: i64=0 70 while k2<K { let w: i64=__f32_from_i64(((j+k2)%4)+1); sb[j*K+k2]=w; pack4(pb,j*K+k2,w); k2=k2+1 } 71 j=j+1 } 72 73 mm_scalar(sa, sb, c0, M, N, K) 74 mm_dot(pa, pb, cd, M, N, K) 75 mm_fma(pa, pb, cf, acc, M, N, K) 76 77 var mmd: i64=0 78 var mmf: i64=0 79 i=0 80 while i<M { var jj: i64=0 81 while jj<N { if __f32_to_i64(cd[i*N+jj]) != __f32_to_i64(c0[i*N+jj]) { mmd=mmd+1 } if __f32_to_i64(cf[i*N+jj]) != __f32_to_i64(c0[i*N+jj]) { mmf=mmf+1 } jj=jj+1 } 82 i=i+1 } 83 let s0: i64=__f32_to_i64(c0[0]) 84 let sf: i64=__f32_to_i64(cf[0]) 85 var exp: i64=0 86 var kk: i64=0 87 while kk<K { let a: i64=((0+kk)%4)+1; exp=exp+a*a; kk=kk+1 } 88 89 let REPS: i64=50 90 let t0: i64=sys_now_us() 91 var r: i64=0 92 while r<REPS { mm_scalar(sa, sb, c0, M, N, K); r=r+1 } 93 let t1: i64=sys_now_us() 94 var rd: i64=0 95 while rd<REPS { mm_dot(pa, pb, cd, M, N, K); rd=rd+1 } 96 let t2: i64=sys_now_us() 97 var rf: i64=0 98 while rf<REPS { mm_fma(pa, pb, cf, acc, M, N, K); rf=rf+1 } 99 let t3: i64=sys_now_us() 100 var usd: i64=t2-t1 101 if usd<=0 { usd=1 } 102 var usf: i64=t3-t2 103 if usf<=0 { usf=1 } 104 var us0: i64=t1-t0 105 if us0<=0 { us0=1 } 106 let spd: i64=us0*100/usd 107 let spf: i64=us0*100/usf 108 let spfd: i64=usd*100/usf 109 110 gx_puts(" C[0][0]: scalar="); gx_num(s0); gx_puts(" fma="); gx_num(sf); gx_puts(" (expected "); gx_num(exp); gx_puts(")\n"); 111 gx_puts(" bit-exact: dot mism="); gx_num(mmd); gx_puts(" fma mism="); gx_num(mmf); gx_puts(" / "); gx_num(M*N); gx_puts("\n"); 112 gx_puts(" scalar="); gx_num(us0); gx_puts("us AVX2-dot="); gx_num(usd); gx_puts("us AVX2-FMA="); gx_num(usf); gx_puts("us\n"); 113 gx_puts(" AVX2-dot = "); gx_num(spd/100); gx_puts("."); gx_num((spd%100)/10); gx_num(spd%10); gx_puts("x over scalar AVX2-FMA = "); gx_num(spf/100); gx_puts("."); gx_num((spf%100)/10); gx_num(spf%10); gx_puts("x over scalar FMA vs dot = "); gx_num(spfd/100); gx_puts("."); gx_num((spfd%100)/10); gx_num(spfd%10); gx_puts("x\n\n"); 114 115 var pass: i64=0 116 var ttl: i64=0 117 ttl=ttl+1; gx_puts(" T1 __f32x8_fma + __f32x8_hsum compiled+ran (FMA encoder via .byte VEX): "); if sf>0 { pass=pass+1; gx_puts("PASS\n") } else { gx_puts("FAIL\n") } 118 ttl=ttl+1; gx_puts(" T2 FMA CORRECT == scalar BIT-EXACT (0 mismatches): "); if mmf==0 { pass=pass+1; gx_puts("PASS\n") } else { gx_puts("FAIL\n") } 119 ttl=ttl+1; gx_puts(" T3 known-answer C[0][0] == "); gx_num(exp); gx_puts(": "); if sf==exp { pass=pass+1; gx_puts("PASS\n") } else { gx_puts("FAIL\n") } 120 ttl=ttl+1; gx_puts(" T4 FMA accumulate FASTER than per-chunk-hsum dot (the fused-MAC + single-hsum win): "); if usf<usd { pass=pass+1; gx_puts("PASS\n") } else { gx_puts("FAIL\n") } 121 122 gx_puts("NX-FMA-MATMUL-GATE passed "); gx_num(pass); gx_puts("/"); gx_num(ttl) 123 // MIGRATED onto nx_gate_verdict by nx_gate_dry_apply (D001, minimal form): every check 124 // row above is untouched, so the PASS/FAIL vector cannot change; only the hand-rolled 125 // verdict emission is replaced by the ONE shared base class. Proven by nx_gate_migrate verify. 126 let ctr__dry: *i64 = gv_ctr() 127 ctr__dry[0] = pass 128 ctr__dry[1] = ttl 129 let rc__dry: i64 = gv_verdict("FMA-MATMUL-GATE" as *u8, ctr__dry, "FMA vector-accumulate WORKS bit-exact + faster than per-chunk AVX2 -- the compute lever)" as *u8) 130 sys_exit(rc__dry) 131 return rc__dry 132}