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}