code wiki / (root) / nx_f32x8_matmul_gate.nx

nx_f32x8_matmul_gate.nx source

↩ module page · 128 lines · 6505 B

1// nx_f32x8_matmul_gate.nx -- PROVES the AVX2 8-wide lever past SSE 4-wide. 2// Three matmuls, same data: scalar (__f32_mul, 1 MAC/instr) vs SSE packed (__f32x4_dot, 4) vs AVX2 packed 3// (__f32x8_dot, 8 -- vmovups+vmulps 256-bit via .byte VEX, proven nxasm_vex_kat 7/7). Verifies all three 4// agree BIT-EXACT (small-int f32 -> order-independent) + a known-answer dot + MEASURES both packed speedups. 5// The AVX2 rung is the gating compute lever for racing sdcpp. Honest: vextractf128+hsum per 8 is overhead; 6// a vector-accumulating FMA kernel approaches the full 8x -- this first AVX2 kernel measures what it measures. 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_x4(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, __f32x4_dot((pab+(i*K+k)*4) as *u8, (pbb+(j*K+k)*4) as *u8)); k=k+4 } 31 c[i*N+j]=acc; j=j+1 } 32 i=i+1 } 33 return 0 34} 35func mm_x8(pa: *u8, pb: *u8, c: *i64, M: i64, N: i64, K: i64) -> i64 { 36 let pab: i64=pa as i64 37 let pbb: i64=pb as i64 38 var i: i64=0 39 while i<M { var j: i64=0 40 while j<N { var acc: i64=__f32_from_i64(0); var k: i64=0 41 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 } 42 c[i*N+j]=acc; j=j+1 } 43 i=i+1 } 44 return 0 45} 46 47func main() -> i64 { 48 gx_puts("AVX2 8-wide matmul: __f32x8_dot (vmovups+vmulps 256-bit) vs SSE __f32x4_dot vs scalar\n\n" as *u8) 49 let M: i64=64 50 let N: i64=64 51 let K: i64=64 52 let sa: *i64 = sys_mmap(M*K*8) as *i64 53 let sb: *i64 = sys_mmap(N*K*8) as *i64 54 let pa: *u8 = sys_mmap(M*K*4) 55 let pb: *u8 = sys_mmap(N*K*4) 56 let c0: *i64 = sys_mmap(M*N*8) as *i64 57 let c4: *i64 = sys_mmap(M*N*8) as *i64 58 let c8: *i64 = sys_mmap(M*N*8) as *i64 59 60 var i: i64=0 61 while i<M { var k: i64=0 62 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 } 63 i=i+1 } 64 var j: i64=0 65 while j<N { var k2: i64=0 66 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 } 67 j=j+1 } 68 69 mm_scalar(sa, sb, c0, M, N, K) 70 mm_x4(pa, pb, c4, M, N, K) 71 mm_x8(pa, pb, c8, M, N, K) 72 73 var mm4: i64=0 74 var mm8: i64=0 75 i=0 76 while i<M { var jj: i64=0 77 while jj<N { if __f32_to_i64(c4[i*N+jj]) != __f32_to_i64(c0[i*N+jj]) { mm4=mm4+1 } if __f32_to_i64(c8[i*N+jj]) != __f32_to_i64(c0[i*N+jj]) { mm8=mm8+1 } jj=jj+1 } 78 i=i+1 } 79 let s0: i64=__f32_to_i64(c0[0]) 80 let s8: i64=__f32_to_i64(c8[0]) 81 var exp: i64=0 82 var kk: i64=0 83 while kk<K { let a: i64=((0+kk)%4)+1; exp=exp+a*a; kk=kk+1 } 84 85 let REPS: i64=50 86 let t0: i64=sys_now_us() 87 var r: i64=0 88 while r<REPS { mm_scalar(sa, sb, c0, M, N, K); r=r+1 } 89 let t1: i64=sys_now_us() 90 var r4: i64=0 91 while r4<REPS { mm_x4(pa, pb, c4, M, N, K); r4=r4+1 } 92 let t2: i64=sys_now_us() 93 var r8: i64=0 94 while r8<REPS { mm_x8(pa, pb, c8, M, N, K); r8=r8+1 } 95 let t3: i64=sys_now_us() 96 var us0: i64=t1-t0 97 if us0<=0 { us0=1 } 98 var us4: i64=t2-t1 99 if us4<=0 { us4=1 } 100 var us8: i64=t3-t2 101 if us8<=0 { us8=1 } 102 let sp4: i64=us0*100/us4 103 let sp8: i64=us0*100/us8 104 let sp84: i64=us4*100/us8 105 106 gx_puts(" C[0][0]: scalar="); gx_num(s0); gx_puts(" avx2="); gx_num(s8); gx_puts(" (expected "); gx_num(exp); gx_puts(")\n"); 107 gx_puts(" bit-exact: x4 mismatches="); gx_num(mm4); gx_puts(" x8 mismatches="); gx_num(mm8); gx_puts(" / "); gx_num(M*N); gx_puts("\n"); 108 gx_puts(" scalar="); gx_num(us0); gx_puts("us SSE-x4="); gx_num(us4); gx_puts("us AVX2-x8="); gx_num(us8); gx_puts("us\n"); 109 gx_puts(" SSE-x4 = "); gx_num(sp4/100); gx_puts("."); gx_num((sp4%100)/10); gx_num(sp4%10); gx_puts("x over scalar AVX2-x8 = "); gx_num(sp8/100); gx_puts("."); gx_num((sp8%100)/10); gx_num(sp8%10); gx_puts("x over scalar AVX2 vs SSE = "); gx_num(sp84/100); gx_puts("."); gx_num((sp84%100)/10); gx_num(sp84%10); gx_puts("x\n\n"); 110 111 var pass: i64=0 112 var ttl: i64=0 113 ttl=ttl+1; gx_puts(" T1 __f32x8_dot (AVX2 256-bit via .byte VEX) compiled+ran: "); if s8>0 { pass=pass+1; gx_puts("PASS\n") } else { gx_puts("FAIL\n") } 114 ttl=ttl+1; gx_puts(" T2 AVX2 CORRECT == scalar BIT-EXACT (0 mismatches): "); if mm8==0 { pass=pass+1; gx_puts("PASS\n") } else { gx_puts("FAIL\n") } 115 ttl=ttl+1; gx_puts(" T3 known-answer C[0][0] == "); gx_num(exp); gx_puts(": "); if s8==exp { pass=pass+1; gx_puts("PASS\n") } else { gx_puts("FAIL\n") } 116 ttl=ttl+1; gx_puts(" T4 AVX2 8-wide FASTER than SSE 4-wide (the 2x-width lever): "); if us8<us4 { pass=pass+1; gx_puts("PASS\n") } else { gx_puts("FAIL\n") } 117 118 gx_puts("NX-F32X8-MATMUL-GATE passed "); gx_num(pass); gx_puts("/"); gx_num(ttl) 119 // MIGRATED onto nx_gate_verdict by nx_gate_dry_apply (D001, minimal form): every check 120 // row above is untouched, so the PASS/FAIL vector cannot change; only the hand-rolled 121 // verdict emission is replaced by the ONE shared base class. Proven by nx_gate_migrate verify. 122 let ctr__dry: *i64 = gv_ctr() 123 ctr__dry[0] = pass 124 ctr__dry[1] = ttl 125 let rc__dry: i64 = gv_verdict("F32X8-MATMUL-GATE" as *u8, ctr__dry, "AVX2 8-wide WORKS end-to-end, bit-exact + faster than SSE -- the next compute rung is REAL)" as *u8) 126 sys_exit(rc__dry) 127 return rc__dry 128}