code wiki / (root) / nx_fma_unroll_matmul_gate.nx

nx_fma_unroll_matmul_gate.nx source

↩ module page · 150 lines · 7511 B

1// nx_fma_unroll_matmul_gate.nx -- unlock the FMA throughput win: MULTIPLE independent accumulators. 2// The single-accumulator FMA (nx_fma_matmul_gate) was only ~1.03x over AVX2-dot because each vfmadd231ps 3// depends on the previous (store->load->fma serial chain) = LATENCY-bound. FMA has ~4-5 cycle latency but 4// ~0.5 cycle throughput, so to fill the pipeline you need >=4 INDEPENDENT accumulators in flight. This 5// unrolls the k-loop by 4 (a0..a3, each its own 8-wide accumulator, disjoint k-chunks) so the 4 FMAs have 6// NO inter-dependency -> the CPU pipelines them -> hides the latency. NishiLang-only (reuses __f32x8_fma/ 7// __f32x8_hsum) -- NO compiler change. Same data, BIT-EXACT vs scalar. No hw writes (Rule 26). 8// expect_exit: 0 license_tier: ORIGINAL 9import "nx_syscalls.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} 24// single accumulator (latency-bound serial chain) 25func mm_fma1(pa: *u8, pb: *u8, c: *i64, acc: *u8, M: i64, N: i64, K: i64) -> i64 { 26 let pab: i64=pa as i64 27 let pbb: i64=pb as i64 28 let az: *i64 = acc as *i64 29 var i: i64=0 30 while i<M { var j: i64=0 31 while j<N { az[0]=0; az[1]=0; az[2]=0; az[3]=0; var k: i64=0 32 while k<K { __f32x8_fma(acc, (pab+(i*K+k)*4) as *u8, (pbb+(j*K+k)*4) as *u8); k=k+8 } 33 c[i*N+j]=__f32x8_hsum(acc); j=j+1 } 34 i=i+1 } 35 return 0 36} 37// FOUR independent accumulators (ILP hides FMA latency) 38func mm_fma4(pa: *u8, pb: *u8, c: *i64, a0: *u8, a1: *u8, a2: *u8, a3: *u8, M: i64, N: i64, K: i64) -> i64 { 39 let pab: i64=pa as i64 40 let pbb: i64=pb as i64 41 let z0: *i64=a0 as *i64 42 let z1: *i64=a1 as *i64 43 let z2: *i64=a2 as *i64 44 let z3: *i64=a3 as *i64 45 var i: i64=0 46 while i<M { var j: i64=0 47 while j<N { 48 z0[0]=0; z0[1]=0; z0[2]=0; z0[3]=0 49 z1[0]=0; z1[1]=0; z1[2]=0; z1[3]=0 50 z2[0]=0; z2[1]=0; z2[2]=0; z2[3]=0 51 z3[0]=0; z3[1]=0; z3[2]=0; z3[3]=0 52 var k: i64=0 53 while k<K { 54 let ba: i64=pab+(i*K+k)*4 55 let bb: i64=pbb+(j*K+k)*4 56 __f32x8_fma(a0, (ba) as *u8, (bb) as *u8) 57 __f32x8_fma(a1, (ba+32) as *u8, (bb+32) as *u8) 58 __f32x8_fma(a2, (ba+64) as *u8, (bb+64) as *u8) 59 __f32x8_fma(a3, (ba+96) as *u8, (bb+96) as *u8) 60 k=k+32 61 } 62 let lo: i64=__f32_add(__f32x8_hsum(a0), __f32x8_hsum(a1)) 63 let hi: i64=__f32_add(__f32x8_hsum(a2), __f32x8_hsum(a3)) 64 c[i*N+j]=__f32_add(lo, hi); j=j+1 } 65 i=i+1 } 66 return 0 67} 68 69func main() -> i64 { 70 gx_puts("FMA throughput: 4 INDEPENDENT accumulators (ILP) vs 1 (latency-bound) vs scalar\n\n" as *u8) 71 let M: i64=64 72 let N: i64=64 73 let K: i64=256 74 let sa: *i64 = sys_mmap(M*K*8) as *i64 75 let sb: *i64 = sys_mmap(N*K*8) as *i64 76 let pa: *u8 = sys_mmap(M*K*4) 77 let pb: *u8 = sys_mmap(N*K*4) 78 let acc: *u8 = sys_mmap(64) 79 let a0: *u8 = sys_mmap(64) 80 let a1: *u8 = sys_mmap(64) 81 let a2: *u8 = sys_mmap(64) 82 let a3: *u8 = sys_mmap(64) 83 let c0: *i64 = sys_mmap(M*N*8) as *i64 84 let c1: *i64 = sys_mmap(M*N*8) as *i64 85 let c4: *i64 = sys_mmap(M*N*8) as *i64 86 87 var i: i64=0 88 while i<M { var k: i64=0 89 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 } 90 i=i+1 } 91 var j: i64=0 92 while j<N { var k2: i64=0 93 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 } 94 j=j+1 } 95 96 mm_scalar(sa, sb, c0, M, N, K) 97 mm_fma1(pa, pb, c1, acc, M, N, K) 98 mm_fma4(pa, pb, c4, a0, a1, a2, a3, M, N, K) 99 100 var m1: i64=0 101 var m4: i64=0 102 i=0 103 while i<M { var jj: i64=0 104 while jj<N { if __f32_to_i64(c1[i*N+jj]) != __f32_to_i64(c0[i*N+jj]) { m1=m1+1 } if __f32_to_i64(c4[i*N+jj]) != __f32_to_i64(c0[i*N+jj]) { m4=m4+1 } jj=jj+1 } 105 i=i+1 } 106 let s4: i64=__f32_to_i64(c4[0]) 107 var exp: i64=0 108 var kk: i64=0 109 while kk<K { let a: i64=((0+kk)%4)+1; exp=exp+a*a; kk=kk+1 } 110 111 let REPS: i64=50 112 let t0: i64=sys_now_us() 113 var r: i64=0 114 while r<REPS { mm_scalar(sa, sb, c0, M, N, K); r=r+1 } 115 let t1: i64=sys_now_us() 116 var r1: i64=0 117 while r1<REPS { mm_fma1(pa, pb, c1, acc, M, N, K); r1=r1+1 } 118 let t2: i64=sys_now_us() 119 var r4: i64=0 120 while r4<REPS { mm_fma4(pa, pb, c4, a0, a1, a2, a3, M, N, K); r4=r4+1 } 121 let t3: i64=sys_now_us() 122 var us0: i64=t1-t0 123 if us0<=0 { us0=1 } 124 var us1: i64=t2-t1 125 if us1<=0 { us1=1 } 126 var us4: i64=t3-t2 127 if us4<=0 { us4=1 } 128 let sp1: i64=us0*100/us1 129 let sp4: i64=us0*100/us4 130 let sp41: i64=us1*100/us4 131 let flop: i64=2*M*N*K 132 let mf4: i64=flop/(us4/REPS+1) 133 let gap4: i64=1842560/(mf4+1) 134 135 gx_puts(" C[0][0] fma4="); gx_num(s4); gx_puts(" (expected "); gx_num(exp); gx_puts(") bit-exact: fma1="); gx_num(m1); gx_puts(" fma4="); gx_num(m4); gx_puts(" / "); gx_num(M*N); gx_puts("\n"); 136 gx_puts(" scalar="); gx_num(us0); gx_puts("us FMA-1acc="); gx_num(us1); gx_puts("us FMA-4acc="); gx_num(us4); gx_puts("us\n"); 137 gx_puts(" FMA-1acc = "); gx_num(sp1/100); gx_puts("."); gx_num((sp1%100)/10); gx_num(sp1%10); gx_puts("x over scalar FMA-4acc = "); gx_num(sp4/100); gx_puts("."); gx_num((sp4%100)/10); gx_num(sp4%10); gx_puts("x over scalar\n"); 138 gx_puts(" 4acc vs 1acc (ILP latency-hiding win) = "); gx_num(sp41/100); gx_puts("."); gx_num((sp41%100)/10); gx_num(sp41%10); gx_puts("x | FMA-4acc single-thread = "); gx_num(mf4); gx_puts(" MFLOP/s (gap to peak ~"); gx_num(gap4); gx_puts("x)\n\n"); 139 140 var pass: i64=0 141 var ttl: i64=0 142 ttl=ttl+1; gx_puts(" T1 4-accumulator FMA BIT-EXACT vs scalar (0 mismatches): "); if m4==0 { pass=pass+1; gx_puts("PASS\n") } else { gx_puts("FAIL\n") } 143 ttl=ttl+1; gx_puts(" T2 known-answer C[0][0] == "); gx_num(exp); gx_puts(": "); if s4==exp { pass=pass+1; gx_puts("PASS\n") } else { gx_puts("FAIL\n") } 144 ttl=ttl+1; gx_puts(" T3 4 accumulators FASTER than 1 (ILP hides FMA latency -- the real throughput win): "); if us4<us1 { pass=pass+1; gx_puts("PASS\n") } else { gx_puts("FAIL\n") } 145 ttl=ttl+1; gx_puts(" T4 FMA-4acc is the fastest single-thread kernel built (>= 9x over scalar): "); if sp4>=900 { pass=pass+1; gx_puts("PASS\n") } else { gx_puts("FAIL\n") } 146 147 gx_puts("NX-FMA-UNROLL-GATE passed "); gx_num(pass); gx_puts("/"); gx_num(ttl) 148 if pass==ttl { gx_puts(" verdict=GREEN (multi-accumulator FMA unlocks the throughput win -- the deepest single-thread compute rung)\n"); sys_exit(0); return 0 } 149 gx_puts(" verdict=RED\n"); sys_exit(1); return 1 150}