code wiki / (root) / nx_f32_matmul_t_test.nx

nx_f32_matmul_t_test.nx source

↩ module page · 45 lines · 1403 B

1// nx_f32_matmul_t_test.nx -- KAT for column-major-B matmul. 2// A=[1,2] @ B (3x2 col-major) -> C[1,2] 3 4import "nx_syscalls.nx" 5import "nx_tier.nx" 6import "nx_f32.nx" 7import "nx_f32_matmul.nx" 8import "nx_f32_matmul_t.nx" 9 10func main() -> i64 { 11 // A = [1.0, 2.0] shape [1, 2] (m=1, k=2) 12 let A: *i64 = sys_mmap(2 * 8) as *i64 13 A[0] = 0x3F800000 // 1.0 14 A[1] = 0x40000000 // 2.0 15 16 // B logical [k=2, n=3]: 17 // [[1, 2, 3], 18 // [4, 5, 6]] 19 // Expected C = A @ B = [1*1+2*4, 1*2+2*5, 1*3+2*6] = [9, 12, 15] 20 // 21 // Column-major storage (dim_0=k=2 fast-varying): 22 // memory[0..1] = column 0 = [1, 4] 23 // memory[2..3] = column 1 = [2, 5] 24 // memory[4..5] = column 2 = [3, 6] 25 let B: *i64 = sys_mmap(6 * 8) as *i64 26 B[0] = 0x3F800000 // 1.0 (col 0, row 0) 27 B[1] = 0x40800000 // 4.0 (col 0, row 1) 28 B[2] = 0x40000000 // 2.0 (col 1, row 0) 29 B[3] = 0x40A00000 // 5.0 (col 1, row 1) 30 B[4] = 0x40400000 // 3.0 (col 2, row 0) 31 B[5] = 0x40C00000 // 6.0 (col 2, row 1) 32 33 let C: *i64 = sys_mmap(3 * 8) as *i64 34 35 nx_f32_matmul_t(A, B, C, 1, 2, 3) 36 37 // C[0] = 1*1 + 2*4 = 9.0 = 0x41100000 38 // C[1] = 1*2 + 2*5 = 12.0 = 0x41400000 39 // C[2] = 1*3 + 2*6 = 15.0 = 0x41700000 40 if C[0] != 0x41100000 { return 10 } 41 if C[1] != 0x41400000 { return 11 } 42 if C[2] != 0x41700000 { return 12 } 43 44 return 0 45}