code wiki / (root) / nx_i8fma32_kat.nx

nx_i8fma32_kat.nx source

↩ module page · 54 lines · 2017 B

1// nx_i8fma32_kat.nx -- KAT for __f32_i8fma32 (deferred-hsum FMA block). 2// acc[8] += d * (sext(a) . b), 8-lane; caller hsums with __f32x8_hsum. 3// Exact-int inputs -> order-independent -> hsum(acc) must EQUAL d*dot and 4// two accumulated blocks must equal the running sum. expect_exit: 0. 5import "nx_syscalls.nx" 6import "nx_f32.nx" 7import "nx_f32_cvt.nx" 8 9func st4(p: *u8, idx: i64, bits: i64) -> i64 { 10 p[idx*4+0]=bits as u8; p[idx*4+1]=(bits>>8) as u8; p[idx*4+2]=(bits>>16) as u8; p[idx*4+3]=(bits>>24) as u8; return 0 11} 12// scalar reference: d * sum(sext(a[j]) * b[j]) 13func ref_block(a: *u8, b: *u8, d: i64) -> i64 { 14 var s: i64 = 0 15 var j: i64 = 0 16 while j < 32 { 17 var av: i64 = a[j] as i64 18 if av >= 128 { av = av - 256 } 19 let bj: i64 = (b[j*4] as i64) | ((b[j*4+1] as i64)<<8) | ((b[j*4+2] as i64)<<16) | ((b[j*4+3] as i64)<<24) 20 s = __f32_add(s, __f32_mul(nx_i32_to_f32(av), bj)) 21 j = j + 1 22 } 23 return __f32_mul(d, s) 24} 25 26func main() -> i64 { 27 let a: *u8 = sys_mmap(32) 28 let b: *u8 = sys_mmap(32*4) 29 let acc: *u8 = sys_mmap(32) // 8 f32 lanes 30 var j: i64 = 0 31 while j < 32 { a[j] = (j-16) as u8; st4(b, j, nx_i32_to_f32((j%5)+1)); j = j+1 } 32 let d: i64 = nx_i32_to_f32(3) 33 34 // one block: acc=0, fma, hsum == d*dot 35 var z: i64 = 0 36 while z < 8 { st4(acc, z, 0); z = z + 1 } 37 __f32_i8fma32(a, b, d, acc as *u8) 38 let got: i64 = __f32x8_hsum(acc as *i64) 39 let ref: i64 = ref_block(a, b, d) 40 if got != ref { return 77 } 41 42 // two blocks accumulate: acc already holds block1; add block2 (a2,b2,d2) 43 let a2: *u8 = sys_mmap(32) 44 let b2: *u8 = sys_mmap(32*4) 45 var j2: i64 = 0 46 while j2 < 32 { a2[j2] = (7 - (j2%9)) as u8; st4(b2, j2, nx_i32_to_f32((j2%4)+2)); j2 = j2+1 } 47 let d2: i64 = nx_i32_to_f32(2) 48 __f32_i8fma32(a2, b2, d2, acc as *u8) 49 let got2: i64 = __f32x8_hsum(acc as *i64) 50 let ref2: i64 = __f32_add(ref, ref_block(a2, b2, d2)) 51 if got2 != ref2 { return 78 } 52 53 return 0 54}