nx_f32_q4k_matmul_test.nx source
↩ module page · 141 lines · 5581 B
1// nx_f32_q4k_matmul_test.nx -- smoke for nx_f32_q4k_matmul.nx.
2//
3// REWRITTEN 2026-07-07 to the organ's CURRENT axis semantics: B is n
4// WEIGHT ROWS, each row = k quantized values ((k/256) super-blocks,
5// 144 B each), C[i,j] = dot(A[i,:], dequant(row j)). The previous
6// version of this test encoded the OLD transposed [in,out] layout
7// (k=1, one block along n) -- after the k-axis fix landed in the
8// organ, that shape trips ERR_ALIGN (k=1 is not 256-aligned) and the
9// test had been failing exit 23 UNNOTICED (found 2026-07-07 during
10// the threading arc's consumer sweep; serving-ne-working class).
11//
12// Builds a synthetic Q4_K super-block whose dequantized values are
13// 2.0 at position 0, 3.0 at position 32, 4.0 at position 128, and 0
14// everywhere else (ggml 32-byte-group nibble layout, see
15// nx_q4k_to_f32_test.nx).
16//
17// Tests (all operands small exact integers -> f32 math is bit-exact):
18// 1. m=1, k=256, n=1: A = 1.0 at [0],[32],[128] -> C[0] = 2+3+4 = 9.0;
19// then A[0]=2.0 -> C[0] = 11.0 (input scaling).
20// 2. n=2 identical rows -> C[0] = C[1] = 9.0 (row offset stride 144).
21// 3. m=2 (second A row doubled) -> row0 gets 9.0s, row1 gets 18.0s.
22// 4. k not 256-aligned -> NX_FQ4M_ERR_ALIGN.
23// 5. MT cross-check: nx_f32_q4k_matmul_mt (2 workers) bit-equals the
24// serial result on the n=2 shape (full adversarial coverage lives
25// in nx_q4k_matmul_mt_gate).
26
27import "nx_syscalls.nx"
28import "nx_tier.nx"
29import "nx_le.nx"
30import "nx_gguf_load.nx"
31import "nx_f32.nx"
32import "nx_q4k_to_f32.nx"
33import "nx_f32_q4k_matmul.nx"
34
35// Write one synthetic super-block (144 bytes) at offset `off` into buf.
36// Layout from nx_q4k_to_f32_test.nx:
37// bytes [0,1]: f16 d = 1.0
38// bytes [2,3]: f16 dmin = 0
39// bytes [4..16]: scales/mins (sc[0..3]=1, m[0..3]=0, sc[4] high-nibble of byte 12 = 2)
40// bytes [16..144]: 128 nibble bytes, all 0 except [16+0]=0x32, [16+64]=0x32
41
42func _write_block(buf: *u8, off: i64) -> i64 {
43 nx_le_write_u16(buf, off + 0, 0x3C00) // d = 1.0
44 nx_le_write_u16(buf, off + 2, 0x0000) // dmin = 0
45 buf[off + 4] = 0x01 as u8
46 buf[off + 5] = 0x01 as u8
47 buf[off + 6] = 0x01 as u8
48 buf[off + 7] = 0x01 as u8
49 buf[off + 8] = 0 as u8
50 buf[off + 9] = 0 as u8
51 buf[off + 10] = 0 as u8
52 buf[off + 11] = 0 as u8
53 buf[off + 12] = 0x02 as u8
54 buf[off + 13] = 0 as u8
55 buf[off + 14] = 0 as u8
56 buf[off + 15] = 0 as u8
57 var z: i64 = 0
58 while z < 128 {
59 buf[off + 16 + z] = 0 as u8
60 z = z + 1
61 }
62 buf[off + 16 + 0] = 0x32 as u8
63 buf[off + 16 + 64] = 0x32 as u8
64 return off + 144
65}
66
67// A-row filler: 1.0 (scaled by `mul_bits` slot value) at the block's
68// three nonzero weight positions, 0 elsewhere.
69func _fill_a_row(A: *i64, row_off: i64, v_bits: i64) -> i64 {
70 var p: i64 = 0
71 while p < 256 {
72 A[row_off + p] = 0
73 p = p + 1
74 }
75 A[row_off + 0] = v_bits
76 A[row_off + 32] = v_bits
77 A[row_off + 128] = v_bits
78 return 0
79}
80
81func main() -> i64 {
82 var vi: nx_int = 0
83 while vi < NX_FQ4M_N_VERDICTS {
84 if nx_fq4m_verdict_is_valid(vi) != 1 { return 5 + vi }
85 vi = vi + 1
86 }
87
88 // ===== Test 1: m=1, k=256, n=1 =====
89 let buf: *u8 = sys_mmap(512)
90 var _end: i64 = _write_block(buf, 0)
91
92 let A1: *i64 = sys_mmap(256 * 8) as *i64
93 _fill_a_row(A1, 0, 0x3F800000) // 1.0 at the 3 hot positions
94 let C1: *i64 = sys_mmap(8) as *i64
95 let v1: nx_int = nx_f32_q4k_matmul(A1, buf, 0, C1, 1, 256, 1)
96 if v1 != NX_FQ4M_OK { return 20 + v1 }
97 if C1[0] != 0x41100000 { return 30 } // 1*2 + 1*3 + 1*4 = 9.0
98
99 A1[0] = 0x40000000 // scale first input to 2.0
100 let v1b: nx_int = nx_f32_q4k_matmul(A1, buf, 0, C1, 1, 256, 1)
101 if v1b != NX_FQ4M_OK { return 20 + v1b }
102 if C1[0] != 0x41300000 { return 34 } // 2*2 + 3 + 4 = 11.0
103 A1[0] = 0x3F800000 // restore
104
105 // ===== Test 2: n=2 identical rows (row stride = 144 bytes) =====
106 let buf2: *u8 = sys_mmap(512)
107 var pos: i64 = 0
108 pos = _write_block(buf2, pos) // row 0
109 pos = _write_block(buf2, pos) // row 1
110 let C2: *i64 = sys_mmap(2 * 8) as *i64
111 let v2: nx_int = nx_f32_q4k_matmul(A1, buf2, 0, C2, 1, 256, 2)
112 if v2 != NX_FQ4M_OK { return 40 + v2 }
113 if C2[0] != 0x41100000 { return 50 }
114 if C2[1] != 0x41100000 { return 51 }
115
116 // ===== Test 3: m=2, second A row doubled =====
117 let A3: *i64 = sys_mmap(2 * 256 * 8) as *i64
118 _fill_a_row(A3, 0, 0x3F800000) // row 0: 1.0s
119 _fill_a_row(A3, 256, 0x40000000) // row 1: 2.0s
120 let C3: *i64 = sys_mmap(2 * 2 * 8) as *i64
121 let v3: nx_int = nx_f32_q4k_matmul(A3, buf2, 0, C3, 2, 256, 2)
122 if v3 != NX_FQ4M_OK { return 60 + v3 }
123 if C3[0] != 0x41100000 { return 70 } // [0,0] = 9.0
124 if C3[1] != 0x41100000 { return 71 } // [0,1] = 9.0
125 if C3[2] != 0x41900000 { return 72 } // [1,0] = 2*(2+3+4) = 18.0
126 if C3[3] != 0x41900000 { return 73 } // [1,1] = 18.0
127
128 // ===== Test 4: bad alignment (k not multiple of 256) =====
129 let v4: nx_int = nx_f32_q4k_matmul(A1, buf, 0, C1, 1, 100, 1)
130 if v4 != NX_FQ4M_ERR_ALIGN { return 80 }
131
132 // ===== Test 5: MT cross-check (2 workers) on the n=2 shape =====
133 let C5: *i64 = sys_mmap(2 * 8) as *i64
134 C5[0] = 0 - 777777; C5[1] = 0 - 777777 // poison
135 let v5: nx_int = nx_f32_q4k_matmul_mt(A1, buf2, 0, C5, 1, 256, 2, 2)
136 if v5 != NX_FQ4M_OK { return 85 }
137 if C5[0] != C2[0] { return 86 }
138 if C5[1] != C2[1] { return 86 }
139
140 return 0
141}