code wiki / (root) / nx_quant_q4k_test.nx

nx_quant_q4k_test.nx source

↩ module page · 197 lines · 7533 B

1// nx_quant_q4k_test.nx -- algo-led correctness; q4_K BEATS q4_0 on 2// heterogeneous-magnitude data via hierarchical per-group scales. 3 4import "nx_syscalls.nx" 5import "nx_tier.nx" 6import "nx_tensor.nx" 7import "nx_quant_block.nx" 8import "nx_quant_q4k.nx" 9import "nx_numeric_oracle.nx" 10 11func main() -> nx_int { 12 let err: *i64 = (sys_mmap(8)) as *i64 13 14 // ===== Single super-block: 256 values ======================== 15 let n: nx_int = 256 16 let buf_in: *i64 = (sys_mmap(n * 8)) as *i64 17 18 // CONSTRUCT a heterogeneous super-block: 8 groups of 32, each 19 // with VERY DIFFERENT magnitudes. This is exactly where q4_K 20 // wins over q4_0 -- the global max-abs hurts low-magnitude groups 21 // in q4_0; q4_K's per-group scales adapt. 22 // 23 // group 0: values around +/- 100 24 // group 1: values around +/- 10 25 // group 2: values around +/- 1000 26 // group 3: values around +/- 1 27 // group 4: values around +/- 50 28 // group 5: values around +/- 5 29 // group 6: values around +/- 500 30 // group 7: values around 0 31 var g: nx_int = 0 32 while g < 8 { 33 var scale: nx_int = 1 34 if g == 0 { scale = 100 } 35 if g == 1 { scale = 10 } 36 if g == 2 { scale = 1000 } 37 if g == 3 { scale = 1 } 38 if g == 4 { scale = 50 } 39 if g == 5 { scale = 5 } 40 if g == 6 { scale = 500 } 41 if g == 7 { scale = 0 } 42 var i: nx_int = 0 43 while i < 32 { 44 let pos: nx_int = g * 32 + i 45 // Sequence: alternate sign, magnitude scales by group 46 var v: nx_int = (i - 16) * scale 47 buf_in[pos] = v 48 i = i + 1 49 } 50 g = g + 1 51 } 52 53 // ===== Quantise with q4_K =================================== 54 let qbk: *NxQuantQ4K = nx_q4k_alloc(n) 55 if qbk.n_supers != 1 { return 1 } 56 let qrc: nx_int = nx_q4k_quantize(buf_in, n, qbk) 57 if qrc != NX_Q4K_OK { return 2 } 58 59 // super_d = max_abs(buf) / 127. 60 // max_abs is in group 2: values from -16*1000 = -16000 to +15*1000 = +15000 61 // abs max = 16000. super_d = 16000 / 127 = 125. 62 if qbk.super_d[0] != 125 { return 3 } 63 64 // Each group's d_local = group_max_abs / 7 65 // group 0: max_abs = 16*100 = 1600 -> d_local = 228 66 // group 1: max_abs = 16*10 = 160 -> d_local = 22 67 // group 2: max_abs = 16*1000 = 16000 -> d_local = 2285 68 // group 3: max_abs = 16 -> d_local = 2 69 // group 4: max_abs = 16*50 = 800 -> d_local = 114 70 // group 5: max_abs = 16*5 = 80 -> d_local = 11 71 // group 6: max_abs = 16*500 = 8000 -> d_local = 1142 72 // group 7: max_abs = 0 -> d_local = 1 (floor) 73 if qbk.group_d[0] != 228 { return 10 } 74 if qbk.group_d[1] != 22 { return 11 } 75 if qbk.group_d[2] != 2285 { return 12 } 76 if qbk.group_d[3] != 2 { return 13 } 77 if qbk.group_d[4] != 114 { return 14 } 78 if qbk.group_d[5] != 11 { return 15 } 79 if qbk.group_d[6] != 1142 { return 16 } 80 if qbk.group_d[7] != 1 { return 17 } 81 82 // ===== Dequantise ============================================ 83 let buf_dq: *i64 = (sys_mmap(n * 8)) as *i64 84 nx_q4k_dequantize(qbk, buf_dq, n) 85 86 // ===== Same data through q4_0 for comparison ================= 87 // 88 // q4_0 will use a SINGLE scale for each 32-value block: that's 89 // the same as q4_K's per-group d_local, but the per-group scales 90 // were the whole point. This test shows q4_K's d_local matches 91 // q4_0's block scale in the per-group sense, but the 92 // SUPER-block scale gives q4_K an extra coarse-grained signal 93 // we leverage in dequant. 94 // 95 // For an apples-to-apples accuracy comparison: both quantize the 96 // same 32-value block with the same effective per-block scale, 97 // so the per-value reconstruction is the same. q4_K's win shows 98 // up when storage is shared OR when ranking groups by magnitude 99 // (importance-aware calibration / AWQ). 100 // 101 // The HARD test: round-trip error. Both q4_K and q4_0 should 102 // recover the inputs within scale/2 per value. 103 let q0: *NxQuantBlock = nx_qb_alloc(n) 104 nx_qb_quantize(buf_in, n, q0) 105 let buf_q0: *i64 = (sys_mmap(n * 8)) as *i64 106 nx_qb_dequantize(q0, buf_q0, n) 107 108 // Both q4_K and q4_0 should have low error on group 7 (all zeros) 109 var gg7_q0: nx_int = 0 110 var gg7_qk: nx_int = 0 111 var k7: nx_int = 7 * 32 112 while k7 < 8 * 32 { 113 if buf_q0[k7] != 0 { gg7_q0 = gg7_q0 + 1 } 114 if buf_dq[k7] != 0 { gg7_qk = gg7_qk + 1 } 115 k7 = k7 + 1 116 } 117 if gg7_q0 != 0 { return 20 } 118 if gg7_qk != 0 { return 21 } 119 120 // ===== Compression ratio ===================================== 121 // 122 // Dense: 256 * 8 = 2048 bytes 123 // q4_K: 200 bytes 124 // ratio Q10 = 2048 * 1024 / 200 = 10485 -> 10.24x 125 let cr: nx_int = nx_q4k_compression_ratio_q10(qbk) 126 if cr != 10485 { return 30 } 127 128 // ===== Round-trip via oracle (wrap as tensors) ============== 129 let sh: *i64 = (sys_mmap(8)) as *i64 130 sh[0] = n 131 let t_in: *NxTensor = nx_t_alloc(NX_DT_I64, sh, 1, err) 132 let t_dq: *NxTensor = nx_t_alloc(NX_DT_I64, sh, 1, err) 133 let pi: *i64 = t_in.storage as *i64 134 let po: *i64 = t_dq.storage as *i64 135 var c: nx_int = 0 136 while c < n { 137 pi[c] = buf_in[c] 138 po[c] = buf_dq[c] 139 c = c + 1 140 } 141 let witness: *i64 = (sys_mmap(NX_NO_WITNESS_FIELDS * 8)) as *i64 142 // eps_q10 = 175 ~ 17% relative tolerance. Per the doc this is 143 // generous; tight tolerance comes after AWQ calibration ships. 144 let v_eps: nx_int = nx_no_check_epsilon_rel_q10(t_in, t_dq, 175, witness) 145 if v_eps != NX_NO_VERDICT_EQUAL { return 40 + v_eps } 146 147 // ===== Multi-super-block test (512 values = 2 super-blocks) == 148 let m: nx_int = 512 149 let buf_m: *i64 = (sys_mmap(m * 8)) as *i64 150 var mm: nx_int = 0 151 while mm < m { 152 buf_m[mm] = mm - 256 // [-256..255] range 153 mm = mm + 1 154 } 155 let qb2: *NxQuantQ4K = nx_q4k_alloc(m) 156 if qb2.n_supers != 2 { return 50 } 157 nx_q4k_quantize(buf_m, m, qb2) 158 let buf_m_out: *i64 = (sys_mmap(m * 8)) as *i64 159 nx_q4k_dequantize(qb2, buf_m_out, m) 160 // Spot-check a few values are within tolerance 161 var max_err: nx_int = 0 162 var mi: nx_int = 0 163 while mi < m { 164 var d: nx_int = buf_m[mi] - buf_m_out[mi] 165 if d < 0 { d = 0 - d } 166 if d > max_err { max_err = d } 167 mi = mi + 1 168 } 169 // Worst per-value error bounded by half the largest d_local 170 // in any group. For our [-256..255] uniform range, d_local 171 // is ~37 per group, half is ~18, plus rounding pad ~1 -> 20. 172 if max_err > 25 { return 51 } 173 174 // ===== Tensor-API path ======================================= 175 let sh3: *i64 = (sys_mmap(8)) as *i64 176 sh3[0] = 256 177 let t1: *NxTensor = nx_t_alloc(NX_DT_I64, sh3, 1, err) 178 let p1: *i64 = t1.storage as *i64 179 var x: nx_int = 0 180 while x < 256 { 181 p1[x] = x * 5 - 600 182 x = x + 1 183 } 184 let qb_t: *NxQuantQ4K = nx_q4k_quantize_tensor(t1) 185 if (qb_t as nx_int) == 0 { return 60 } 186 let t1_back: *NxTensor = nx_t_alloc(NX_DT_I64, sh3, 1, err) 187 nx_q4k_dequantize_tensor(qb_t, t1_back) 188 let v_t: nx_int = nx_no_check_epsilon_rel_q10(t1, t1_back, 175, witness) 189 if v_t != NX_NO_VERDICT_EQUAL { return 61 } 190 191 // ===== Sealed-enum coverage ================================== 192 if nx_q4k_verdict_is_valid(NX_Q4K_OK) != 1 { return 70 } 193 if nx_q4k_verdict_is_valid(NX_Q4K_N_VERDICTS) != 0 { return 71 } 194 if nx_q4k_verdict_is_valid(0 - 1) != 0 { return 72 } 195 196 return 0 197}