IQ2_XS
Browse files- hexstate_quantize.c +761 -123
- hexstate_requantize.py +134 -20
hexstate_quantize.c
CHANGED
|
@@ -1793,11 +1793,13 @@ void hexstate_set_spectral_params(float dc_lambda, float vw_lambda, float dc_dec
|
|
| 1793 |
g_hex_dc_decay = dc_decay;
|
| 1794 |
}
|
| 1795 |
|
| 1796 |
-
/* Spectral penalty:
|
| 1797 |
-
|
|
|
|
|
|
|
| 1798 |
{
|
| 1799 |
if (HEX_DC_LAMBDA == 0.0f && HEX_VW_LAMBDA == 0.0f) return 0.0f;
|
| 1800 |
-
float dc =
|
| 1801 |
int half = n / 2;
|
| 1802 |
for (int i = 0; i < half; i++) {
|
| 1803 |
float v = e[i] + e[i + half];
|
|
@@ -1808,6 +1810,190 @@ static inline float hex_spectral_penalty(const float *e, int n)
|
|
| 1808 |
+ (HEX_VW_LAMBDA / (float)n) * ves;
|
| 1809 |
}
|
| 1810 |
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1811 |
/* Robust temperature estimator for the HExState measurement model.
|
| 1812 |
*
|
| 1813 |
* The old path estimated T from the mean of each block's MAXIMUM candidate
|
|
@@ -3607,49 +3793,29 @@ static void quantize_tensor_q2k_hpc(const float *weights, int64_t n_elements,
|
|
| 3607 |
}
|
| 3608 |
|
| 3609 |
/* ββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββ
|
| 3610 |
-
* PHASE 3.9 β ROLLING DC
|
| 3611 |
-
*
|
| 3612 |
-
* Transforms the tensor from a collection of isolated 256-element
|
| 3613 |
-
* Q2_K superblocks into a single, continuous error-cancelling waveform.
|
| 3614 |
-
*
|
| 3615 |
-
* After Phase 3 has selected the optimal (d, dmin) candidate for every
|
| 3616 |
-
* block, this sequential pass computes the net DC residual left by each
|
| 3617 |
-
* block using a cheap round-nearest forward quantization, then feeds the
|
| 3618 |
-
* negated, exponentially-decayed residual as a correction bias into the
|
| 3619 |
-
* WLS solver of the immediately following block.
|
| 3620 |
-
*
|
| 3621 |
-
* Mathematically, for block N with final DC residual R_N = Ξ£ Ξ΅α΅’:
|
| 3622 |
-
*
|
| 3623 |
-
* dc_bias[N+1] = βDC_DECAY Γ R_N / QK_K (per-element offset)
|
| 3624 |
*
|
| 3625 |
-
*
|
| 3626 |
-
*
|
|
|
|
| 3627 |
*
|
| 3628 |
-
*
|
|
|
|
| 3629 |
*
|
| 3630 |
-
*
|
| 3631 |
-
*
|
| 3632 |
-
* Rβ, DC_DECAYΒ·Rβ, DC_DECAYΒ²Β·Rβ, β¦ β 0
|
| 3633 |
-
*
|
| 3634 |
-
* The result is written into block_dc_bias[n_blocks]. Phase 4 reads
|
| 3635 |
-
* this array (safe: written sequentially before the parallel loop).
|
| 3636 |
* ββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββ */
|
| 3637 |
|
| 3638 |
#define DC_DECAY (g_hex_dc_decay)
|
| 3639 |
|
| 3640 |
-
float *
|
| 3641 |
|
| 3642 |
-
if (
|
| 3643 |
float rolling_dc = 0.0f;
|
| 3644 |
-
/* Row-boundary awareness: blocks_per_row = row_width / QK_K.
|
| 3645 |
-
* Rows are independent dot products β DC residual from the end of
|
| 3646 |
-
* one row must NOT leak into the start of the next. When row_width
|
| 3647 |
-
* is 0 (unknown), fall back to the old flat-stream behaviour. */
|
| 3648 |
int64_t blocks_per_row = (row_width > 0 && row_width % QK_K == 0)
|
| 3649 |
? row_width / QK_K : 0;
|
| 3650 |
|
| 3651 |
for (int64_t blk = 0; blk < n_blocks; blk++) {
|
| 3652 |
-
/* Reset DC at every row boundary */
|
| 3653 |
if (blocks_per_row > 0 && (blk % blocks_per_row) == 0)
|
| 3654 |
rolling_dc = 0.0f;
|
| 3655 |
|
|
@@ -3664,29 +3830,25 @@ static void quantize_tensor_q2k_hpc(const float *weights, int64_t n_elements,
|
|
| 3664 |
hex_derive_subscales(seeds[blk].scales, seeds[blk].mins,
|
| 3665 |
dm0, mm0, dc_Ls, dc_Lm);
|
| 3666 |
|
| 3667 |
-
/*
|
| 3668 |
-
|
| 3669 |
-
|
| 3670 |
|
| 3671 |
-
/* Quick round-nearest quant to estimate DC residual for NEXT block.
|
| 3672 |
-
* We quantize the adjusted target xβ² = x β dc_bias, then measure
|
| 3673 |
-
* the residual of the ORIGINAL weight against the chosen code. */
|
| 3674 |
float dc_res = 0.0f;
|
| 3675 |
int j, k;
|
| 3676 |
for (j = 0; j < N_SUB; j++) {
|
| 3677 |
float d_sub = dm0 * (float)dc_Ls[j];
|
| 3678 |
float m_sub = mm0 * (float)dc_Lm[j];
|
| 3679 |
for (k = 0; k < 16; k++) {
|
| 3680 |
-
float
|
| 3681 |
int q = 0;
|
| 3682 |
if (d_sub >= 1e-15f) {
|
| 3683 |
-
q = gguf_nearest_int((
|
| 3684 |
if (q < 0) q = 0;
|
| 3685 |
if (q > 3) q = 3;
|
| 3686 |
}
|
| 3687 |
float deq = d_sub * (float)q - m_sub;
|
| 3688 |
-
|
| 3689 |
-
dc_res += bx[16*j + k] - deq;
|
| 3690 |
}
|
| 3691 |
}
|
| 3692 |
rolling_dc = dc_res;
|
|
@@ -3710,21 +3872,7 @@ static void quantize_tensor_q2k_hpc(const float *weights, int64_t n_elements,
|
|
| 3710 |
const float *block_x = weights + blk * QK_K;
|
| 3711 |
int cidx = best_candidate[blk];
|
| 3712 |
uint8_t Ls_blk[16], Lm_blk[16];
|
| 3713 |
-
|
| 3714 |
-
/* ββ Rolling DC boundary condition ββββββββββββββββββββββββββββββ
|
| 3715 |
-
* dc_adj shifts every WLS target in this block so that the net
|
| 3716 |
-
* quantisation error steers toward cancelling the previous block's
|
| 3717 |
-
* DC residual (written by the sequential Phase 3.9 pre-pass). */
|
| 3718 |
-
float dc_adj = (block_dc_bias) ? block_dc_bias[blk] : 0.0f;
|
| 3719 |
-
|
| 3720 |
-
/* Adjusted weight view β WLS and sieve work on this array;
|
| 3721 |
-
* the final error is always reported against the original block_x. */
|
| 3722 |
-
float adj_block_x[QK_K];
|
| 3723 |
-
{
|
| 3724 |
-
int _i;
|
| 3725 |
-
for (_i = 0; _i < QK_K; _i++)
|
| 3726 |
-
adj_block_x[_i] = block_x[_i] - dc_adj;
|
| 3727 |
-
}
|
| 3728 |
|
| 3729 |
uint16_t base_c_d16, base_c_m16;
|
| 3730 |
hex_candidate_pair(seeds[blk].base_dm, seeds[blk].base_mm, cidx, &base_c_d16, &base_c_m16);
|
|
@@ -3742,7 +3890,7 @@ static void quantize_tensor_q2k_hpc(const float *weights, int64_t n_elements,
|
|
| 3742 |
float state_err[N_SUB][6];
|
| 3743 |
|
| 3744 |
for (int j = 0; j < N_SUB; j++) {
|
| 3745 |
-
const float *sx =
|
| 3746 |
for (int v = 0; v < 6; v++) state_err[j][v] = 1e30f;
|
| 3747 |
|
| 3748 |
for (int try_ls = 0; try_ls <= 15; try_ls++) {
|
|
@@ -3850,7 +3998,7 @@ static void quantize_tensor_q2k_hpc(const float *weights, int64_t n_elements,
|
|
| 3850 |
continue;
|
| 3851 |
}
|
| 3852 |
for (int k = 0; k < 16; k++) {
|
| 3853 |
-
int q = gguf_nearest_int((
|
| 3854 |
if (q < 0) q = 0; if (q > 3) q = 3;
|
| 3855 |
L[16*j+k] = (uint8_t)q;
|
| 3856 |
}
|
|
@@ -3861,7 +4009,7 @@ static void quantize_tensor_q2k_hpc(const float *weights, int64_t n_elements,
|
|
| 3861 |
float ls_f = (float)Ls_blk[j];
|
| 3862 |
float lm_f = (float)Lm_blk[j];
|
| 3863 |
for (int k = 0; k < 16; k++) {
|
| 3864 |
-
float x =
|
| 3865 |
float w = (imat_importance) ?
|
| 3866 |
imat_importance[blk * QK_K + 16*j+k] : 1.0f;
|
| 3867 |
float a = ls_f * (float)L[16*j+k];
|
|
@@ -3917,7 +4065,7 @@ static void quantize_tensor_q2k_hpc(const float *weights, int64_t n_elements,
|
|
| 3917 |
float d_sub = trial_dm * (float)Ls_blk[j];
|
| 3918 |
float m_sub = trial_mm * (float)Lm_blk[j];
|
| 3919 |
for (int k = 0; k < 16; k++) {
|
| 3920 |
-
float x =
|
| 3921 |
float w = (imat_importance) ?
|
| 3922 |
imat_importance[blk * QK_K + 16*j+k] : 1.0f;
|
| 3923 |
int q;
|
|
@@ -3943,7 +4091,7 @@ static void quantize_tensor_q2k_hpc(const float *weights, int64_t n_elements,
|
|
| 3943 |
}
|
| 3944 |
|
| 3945 |
for (int j = 0; j < N_SUB; j++) {
|
| 3946 |
-
const float *sx =
|
| 3947 |
float best_sub_err = 1e30f;
|
| 3948 |
uint8_t best_ls = Ls_blk[j], best_lm = Lm_blk[j];
|
| 3949 |
for (int try_ls = 0; try_ls <= 15; try_ls++) {
|
|
@@ -4010,8 +4158,7 @@ static void quantize_tensor_q2k_hpc(const float *weights, int64_t n_elements,
|
|
| 4010 |
q_cont_all[i] = 0.0f;
|
| 4011 |
q_base_all[i] = 0;
|
| 4012 |
} else {
|
| 4013 |
-
|
| 4014 |
-
float qc = (adj_block_x[i] + m_s) / d_s;
|
| 4015 |
q_cont_all[i] = qc;
|
| 4016 |
int qr = gguf_nearest_int(qc);
|
| 4017 |
if (qr < 0) qr = 0; if (qr > 3) qr = 3;
|
|
@@ -4021,17 +4168,19 @@ static void quantize_tensor_q2k_hpc(const float *weights, int64_t n_elements,
|
|
| 4021 |
memcpy(q_shaped_all, q_base_all, QK_K * sizeof(int));
|
| 4022 |
|
| 4023 |
float e_live[QK_K];
|
| 4024 |
-
float dc_cur =
|
|
|
|
| 4025 |
for (int i = 0; i < QK_K; i++) {
|
| 4026 |
int jj = i >> 4;
|
| 4027 |
float d_s = dm * (float)(output[blk].scales[jj] & 0xF);
|
| 4028 |
float m_s = mm * (float)(output[blk].scales[jj] >> 4);
|
| 4029 |
float deq = d_s * (float)q_shaped_all[i] - m_s;
|
| 4030 |
-
/* Vesica on ORIGINAL residuals. DC on the rolling-adjusted
|
| 4031 |
-
* residual so this block still absorbs Phase 3.9's target. */
|
| 4032 |
e_live[i] = block_x[i] - deq;
|
| 4033 |
-
dc_cur +=
|
|
|
|
|
|
|
| 4034 |
}
|
|
|
|
| 4035 |
|
| 4036 |
float v_live[QK_K / 2];
|
| 4037 |
float vesica_cur = 0.0f;
|
|
@@ -4064,6 +4213,9 @@ static void quantize_tensor_q2k_hpc(const float *weights, int64_t n_elements,
|
|
| 4064 |
if (q_try == q_cur) continue;
|
| 4065 |
|
| 4066 |
float e_new = block_x[k] - (d_s * (float)q_try - m_s);
|
|
|
|
|
|
|
|
|
|
| 4067 |
float de = e_new - e_live[k];
|
| 4068 |
|
| 4069 |
float v_new = v_live[pi] + de;
|
|
@@ -4096,29 +4248,35 @@ static void quantize_tensor_q2k_hpc(const float *weights, int64_t n_elements,
|
|
| 4096 |
metric_cur = 4.0f * vesica_cur + dc_cur * dc_cur;
|
| 4097 |
v_live[pi_c] = v_new_c;
|
| 4098 |
e_live[best_k]= e_new_c;
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 4099 |
}
|
| 4100 |
}
|
| 4101 |
|
| 4102 |
-
|
| 4103 |
-
float err_base = 0.0f, err_shaped = 0.0f;
|
| 4104 |
float e_qb[QK_K], e_qs[QK_K];
|
| 4105 |
for (int i = 0; i < QK_K; i++) {
|
| 4106 |
int jj = i >> 4;
|
| 4107 |
float d_s = dm * (float)(output[blk].scales[jj] & 0xF);
|
| 4108 |
float m_s = mm * (float)(output[blk].scales[jj] >> 4);
|
| 4109 |
float w = (imat_importance) ? imat_importance[blk * QK_K + i] : 1.0f;
|
| 4110 |
-
float deq_b = d_s * (float)q_base_all[i] - m_s;
|
| 4111 |
float deq_s = d_s * (float)q_shaped_all[i] - m_s;
|
| 4112 |
-
float xv = block_x[i];
|
| 4113 |
e_qb[i] = xv - deq_b;
|
| 4114 |
e_qs[i] = xv - deq_s;
|
| 4115 |
-
|
| 4116 |
-
|
| 4117 |
}
|
| 4118 |
-
err_base
|
| 4119 |
-
err_shaped
|
| 4120 |
{
|
| 4121 |
-
int use_shaped = (
|
|
|
|
| 4122 |
for (int i = 0; i < QK_K; i++)
|
| 4123 |
L[i] = (uint8_t)(use_shaped ? q_shaped_all[i] : q_base_all[i]);
|
| 4124 |
}
|
|
@@ -4157,7 +4315,7 @@ static void quantize_tensor_q2k_hpc(const float *weights, int64_t n_elements,
|
|
| 4157 |
for (fs_k = 0; fs_k < 16; fs_k++) {
|
| 4158 |
int idx = base + fs_k;
|
| 4159 |
float x_orig = block_x[idx];
|
| 4160 |
-
float x_adj =
|
| 4161 |
|
| 4162 |
/* Propose new code from diffused target */
|
| 4163 |
int q_fs = gguf_nearest_int((x_adj + m_s) / d_s);
|
|
@@ -4178,7 +4336,7 @@ static void quantize_tensor_q2k_hpc(const float *weights, int64_t n_elements,
|
|
| 4178 |
/* Propagate 7/16 of the residual (adj target vs committed code) */
|
| 4179 |
{
|
| 4180 |
float deq_final = d_s * (float)L[idx] - m_s;
|
| 4181 |
-
float residual = (
|
| 4182 |
carry = (fs_k < 15) ? residual * (7.0f / 16.0f) : 0.0f;
|
| 4183 |
}
|
| 4184 |
}
|
|
@@ -4199,7 +4357,7 @@ static void quantize_tensor_q2k_hpc(const float *weights, int64_t n_elements,
|
|
| 4199 |
{
|
| 4200 |
/* Compute current whole-block DC and per-weight residuals */
|
| 4201 |
float wb_e[QK_K];
|
| 4202 |
-
float wb_dc = 0.0f;
|
| 4203 |
for (int i = 0; i < QK_K; i++) {
|
| 4204 |
int jj = i >> 4;
|
| 4205 |
float d_s = dm * (float)(output[blk].scales[jj] & 0xF);
|
|
@@ -4207,21 +4365,21 @@ static void quantize_tensor_q2k_hpc(const float *weights, int64_t n_elements,
|
|
| 4207 |
float deq = d_s * (float)L[i] - m_s;
|
| 4208 |
wb_e[i] = block_x[i] - deq;
|
| 4209 |
wb_dc += wb_e[i];
|
|
|
|
|
|
|
| 4210 |
}
|
| 4211 |
|
| 4212 |
-
|
| 4213 |
-
|
| 4214 |
-
float median_step = dm * 4.0f; /* rough: d * median(Ls) β d*4 */
|
| 4215 |
if (median_step < 1e-15f) median_step = 1e-15f;
|
|
|
|
| 4216 |
|
| 4217 |
for (int dc_pass = 0; dc_pass < 32; dc_pass++) {
|
| 4218 |
-
if (fabsf(
|
| 4219 |
|
| 4220 |
-
/* Find the weight whose nudge toward reducing DC costs the
|
| 4221 |
-
* least in weighted SSE. */
|
| 4222 |
int best_i = -1;
|
| 4223 |
int best_q = 0;
|
| 4224 |
-
float best_ratio = 0.0f;
|
| 4225 |
|
| 4226 |
for (int i = 0; i < QK_K; i++) {
|
| 4227 |
int jj = i >> 4;
|
|
@@ -4229,28 +4387,22 @@ static void quantize_tensor_q2k_hpc(const float *weights, int64_t n_elements,
|
|
| 4229 |
float m_s = mm * (float)(output[blk].scales[jj] >> 4);
|
| 4230 |
if (d_s < 1e-15f) continue;
|
| 4231 |
|
| 4232 |
-
/* Nudge direction: if DC > 0, we need negative residual
|
| 4233 |
-
* change (increase q to reduce xβdeq), and vice versa. */
|
| 4234 |
int q_cur = (int)L[i];
|
| 4235 |
-
int q_try = (
|
| 4236 |
if (q_try < 0 || q_try > 3) continue;
|
| 4237 |
|
| 4238 |
float deq_new = d_s * (float)q_try - m_s;
|
| 4239 |
float e_new = block_x[i] - deq_new;
|
| 4240 |
-
float dc_reduction = fabsf(
|
| 4241 |
if (dc_reduction <= 0.0f) continue;
|
| 4242 |
|
| 4243 |
float w = (imat_importance) ?
|
| 4244 |
imat_importance[blk * QK_K + i] : 1.0f;
|
| 4245 |
-
float
|
| 4246 |
-
|
| 4247 |
-
float sse_cost = sse_new -
|
| 4248 |
-
if (sse_cost < 0.0f) sse_cost = 0.0f;
|
| 4249 |
-
|
| 4250 |
-
/* Ratio: how much DC reduction per unit SSE cost.
|
| 4251 |
-
* Free moves (sse_cost β€ 0) get infinite ratio. */
|
| 4252 |
-
float ratio = (sse_cost < 1e-30f) ? 1e30f
|
| 4253 |
-
: dc_reduction / sse_cost;
|
| 4254 |
if (ratio > best_ratio) {
|
| 4255 |
best_ratio = ratio;
|
| 4256 |
best_i = i;
|
|
@@ -4258,17 +4410,19 @@ static void quantize_tensor_q2k_hpc(const float *weights, int64_t n_elements,
|
|
| 4258 |
}
|
| 4259 |
}
|
| 4260 |
|
| 4261 |
-
if (best_i < 0) break;
|
| 4262 |
|
| 4263 |
-
/* Accept: the minimum-cost nudge toward DC reduction.
|
| 4264 |
-
* Only accept if SSE cost is modest relative to DC gain. */
|
| 4265 |
{
|
| 4266 |
int jj = best_i >> 4;
|
| 4267 |
float d_s = dm * (float)(output[blk].scales[jj] & 0xF);
|
| 4268 |
float m_s = mm * (float)(output[blk].scales[jj] >> 4);
|
| 4269 |
float deq_new = d_s * (float)best_q - m_s;
|
| 4270 |
float e_new = block_x[best_i] - deq_new;
|
| 4271 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
| 4272 |
wb_e[best_i] = e_new;
|
| 4273 |
L[best_i] = (uint8_t)best_q;
|
| 4274 |
}
|
|
@@ -4311,9 +4465,10 @@ static void quantize_tensor_q2k_hpc(const float *weights, int64_t n_elements,
|
|
| 4311 |
* Ξ»_dc/n; vesica/wave handled by the extended-E acceptance. */
|
| 4312 |
{
|
| 4313 |
double rw = (double)HEX_DC_LAMBDA / (double)QK_K;
|
|
|
|
| 4314 |
rSaa += rw * rA * rA; rSab += rw * rA * rB;
|
| 4315 |
-
rSbb += rw * rB * rB; rSxa += rw *
|
| 4316 |
-
rSxb += rw *
|
| 4317 |
}
|
| 4318 |
double rdet = rSaa * rSbb - rSab * rSab;
|
| 4319 |
if (fabs(rdet) > 1e-30) {
|
|
@@ -4343,9 +4498,12 @@ static void quantize_tensor_q2k_hpc(const float *weights, int64_t n_elements,
|
|
| 4343 |
err_try += e_rt[idx] * e_rt[idx] * w;
|
| 4344 |
}
|
| 4345 |
}
|
| 4346 |
-
|
| 4347 |
-
|
| 4348 |
-
|
|
|
|
|
|
|
|
|
|
| 4349 |
}
|
| 4350 |
}
|
| 4351 |
output[blk].d = gguf_fp32_to_fp16(dm);
|
|
@@ -4357,7 +4515,7 @@ static void quantize_tensor_q2k_hpc(const float *weights, int64_t n_elements,
|
|
| 4357 |
* Objective-function mismatch fix: the final passes that commit the
|
| 4358 |
* 2-bit codes β the 16Γ16 (ls, lm) sub-block search, the Β±8 ULP
|
| 4359 |
* (d, dmin) neighborhood search, and the greedy-descent error shaping
|
| 4360 |
-
* β all minimise error against the DC-ADJUSTED target
|
| 4361 |
* The reported RMSE, however, is measured against the ORIGINAL
|
| 4362 |
* weights. The codes are therefore stranded at the optimum of a
|
| 4363 |
* SHIFTED objective, while only the scalar (d, dmin) refit above
|
|
@@ -4443,7 +4601,8 @@ static void quantize_tensor_q2k_hpc(const float *weights, int64_t n_elements,
|
|
| 4443 |
|
| 4444 |
/* Extended score of the CURRENT committed state */
|
| 4445 |
float best_sub = sub_sse[j]
|
| 4446 |
-
+ (HEX_DC_LAMBDA / (float)QK_K)
|
|
|
|
| 4447 |
+ (HEX_VW_LAMBDA / (float)QK_K) * ves_tot;
|
| 4448 |
int best_ls = -1, best_lm = 0;
|
| 4449 |
uint8_t best_q[16];
|
|
@@ -4480,9 +4639,11 @@ static void quantize_tensor_q2k_hpc(const float *weights, int64_t n_elements,
|
|
| 4480 |
if (sub_err >= best_sub) { aborted = 1; break; }
|
| 4481 |
}
|
| 4482 |
if (aborted) continue;
|
|
|
|
|
|
|
| 4483 |
float score = sub_err
|
| 4484 |
+ (HEX_DC_LAMBDA / (float)QK_K)
|
| 4485 |
-
*
|
| 4486 |
+ (HEX_VW_LAMBDA / (float)QK_K)
|
| 4487 |
* (ves_rest + vesc);
|
| 4488 |
if (score < best_sub) {
|
|
@@ -4536,9 +4697,10 @@ static void quantize_tensor_q2k_hpc(const float *weights, int64_t n_elements,
|
|
| 4536 |
}
|
| 4537 |
{
|
| 4538 |
double pw = (double)HEX_DC_LAMBDA / (double)QK_K;
|
|
|
|
| 4539 |
pSaa += pw * pA * pA; pSab += pw * pA * pB;
|
| 4540 |
-
pSbb += pw * pB * pB; pSxa += pw *
|
| 4541 |
-
pSxb += pw *
|
| 4542 |
}
|
| 4543 |
double pdet = pSaa * pSbb - pSab * pSab;
|
| 4544 |
if (fabs(pdet) > 1e-30) {
|
|
@@ -4570,9 +4732,10 @@ static void quantize_tensor_q2k_hpc(const float *weights, int64_t n_elements,
|
|
| 4570 |
err_try += e_pt[idx] * e_pt[idx] * w;
|
| 4571 |
}
|
| 4572 |
}
|
| 4573 |
-
|
| 4574 |
-
|
| 4575 |
-
|
|
|
|
| 4576 |
dm = dm_try;
|
| 4577 |
mm = mm_try;
|
| 4578 |
pol_improved = 1;
|
|
@@ -4609,7 +4772,7 @@ static void quantize_tensor_q2k_hpc(const float *weights, int64_t n_elements,
|
|
| 4609 |
cur_err += e_u[idx] * e_u[idx] * w;
|
| 4610 |
}
|
| 4611 |
}
|
| 4612 |
-
cur_err +=
|
| 4613 |
|
| 4614 |
float best_err = cur_err;
|
| 4615 |
uint16_t best_d16 = base_d16, best_m16 = base_m16;
|
|
@@ -4639,7 +4802,7 @@ static void quantize_tensor_q2k_hpc(const float *weights, int64_t n_elements,
|
|
| 4639 |
if (err >= best_err) { pruned = 1; break; }
|
| 4640 |
}
|
| 4641 |
if (pruned) continue;
|
| 4642 |
-
err +=
|
| 4643 |
if (err < best_err) {
|
| 4644 |
best_err = err;
|
| 4645 |
best_d16 = (uint16_t)cd16;
|
|
@@ -4690,8 +4853,8 @@ static void quantize_tensor_q2k_hpc(const float *weights, int64_t n_elements,
|
|
| 4690 |
fin_err += e_f[idx] * e_f[idx] * w;
|
| 4691 |
}
|
| 4692 |
}
|
| 4693 |
-
|
| 4694 |
-
|
| 4695 |
float g_best = candidate_errors[blk][0];
|
| 4696 |
int g_cand = 0;
|
| 4697 |
for (int c = 1; c < TOTAL_SCALE_CANDIDATES; c++) {
|
|
@@ -4748,11 +4911,53 @@ static void quantize_tensor_q2k_hpc(const float *weights, int64_t n_elements,
|
|
| 4748 |
total_err += berr;
|
| 4749 |
}
|
| 4750 |
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 4751 |
for (int _ti = 0; _ti < _n_omp_threads; _ti++)
|
| 4752 |
hpc_destroy(_tl_graphs[_ti]);
|
| 4753 |
free(_tl_graphs);
|
| 4754 |
|
| 4755 |
-
free(
|
| 4756 |
free(seeds);
|
| 4757 |
free(candidate_errors);
|
| 4758 |
free(best_candidate);
|
|
@@ -5336,6 +5541,418 @@ write_fail:
|
|
| 5336 |
return -1;
|
| 5337 |
}
|
| 5338 |
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 5339 |
/* βββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββ
|
| 5340 |
* LIBRARY API β Exported functions for Python ctypes integration
|
| 5341 |
*
|
|
@@ -5442,6 +6059,27 @@ void hexstate_quantize_tensor_q8_0_hpc(const float *weights, int64_t n_elements,
|
|
| 5442 |
imat_importance, verbose);
|
| 5443 |
}
|
| 5444 |
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 5445 |
#ifndef HEXSTATE_LIBRARY
|
| 5446 |
/* βββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββ
|
| 5447 |
* MAIN
|
|
|
|
| 1793 |
g_hex_dc_decay = dc_decay;
|
| 1794 |
}
|
| 1795 |
|
| 1796 |
+
/* Spectral penalty: (Ξ£e + dc_carry)Β² + Ξ£ vesicaΒ². dc_carry is the
|
| 1797 |
+
* rolling residual we want this block to cancel (Ξ£e β βdc_carry).
|
| 1798 |
+
* Reconstruction target is always the true weight; we never quantize xβbias. */
|
| 1799 |
+
static inline float hex_spectral_penalty_ex(const float *e, int n, float dc_carry)
|
| 1800 |
{
|
| 1801 |
if (HEX_DC_LAMBDA == 0.0f && HEX_VW_LAMBDA == 0.0f) return 0.0f;
|
| 1802 |
+
float dc = dc_carry, ves = 0.0f;
|
| 1803 |
int half = n / 2;
|
| 1804 |
for (int i = 0; i < half; i++) {
|
| 1805 |
float v = e[i] + e[i + half];
|
|
|
|
| 1810 |
+ (HEX_VW_LAMBDA / (float)n) * ves;
|
| 1811 |
}
|
| 1812 |
|
| 1813 |
+
static inline float hex_spectral_penalty(const float *e, int n)
|
| 1814 |
+
{
|
| 1815 |
+
return hex_spectral_penalty_ex(e, n, 0.0f);
|
| 1816 |
+
}
|
| 1817 |
+
|
| 1818 |
+
/* Relative SSE we may spend to hit Ξ£e β βcarry. 5e-4 β ~0.025% RMSE.
|
| 1819 |
+
* Runtime-tunable (hexstate_set_sse_budget) β codebook formats need more. */
|
| 1820 |
+
static float g_hex_sse_budget = 5.0e-4f;
|
| 1821 |
+
#define HEX_DC_SSE_BUDGET (g_hex_sse_budget)
|
| 1822 |
+
void hexstate_set_sse_budget(float rel) { if (rel >= 0.0f) g_hex_sse_budget = rel; }
|
| 1823 |
+
|
| 1824 |
+
static inline void hex_q2k_unpack_L(const BlockQ2K *b, uint8_t L[QK_K])
|
| 1825 |
+
{
|
| 1826 |
+
for (int j = 0; j < QK_K; j += 128) {
|
| 1827 |
+
for (int l = 0; l < 32; l++) {
|
| 1828 |
+
uint8_t p = b->qs[j / 4 + l];
|
| 1829 |
+
L[j + l] = (uint8_t)( p & 3);
|
| 1830 |
+
L[j + l + 32] = (uint8_t)((p >> 2) & 3);
|
| 1831 |
+
L[j + l + 64] = (uint8_t)((p >> 4) & 3);
|
| 1832 |
+
L[j + l + 96] = (uint8_t)((p >> 6) & 3);
|
| 1833 |
+
}
|
| 1834 |
+
}
|
| 1835 |
+
}
|
| 1836 |
+
|
| 1837 |
+
static inline void hex_q2k_pack_L(BlockQ2K *b, const uint8_t L[QK_K])
|
| 1838 |
+
{
|
| 1839 |
+
for (int j = 0; j < QK_K; j += 128) {
|
| 1840 |
+
for (int l = 0; l < 32; l++) {
|
| 1841 |
+
b->qs[j / 4 + l] = (uint8_t)(L[j + l]
|
| 1842 |
+
| (L[j + l + 32] << 2)
|
| 1843 |
+
| (L[j + l + 64] << 4)
|
| 1844 |
+
| (L[j + l + 96] << 6));
|
| 1845 |
+
}
|
| 1846 |
+
}
|
| 1847 |
+
}
|
| 1848 |
+
|
| 1849 |
+
static inline float hex_q2k_el_w(const float *imat, int64_t blk, int i)
|
| 1850 |
+
{
|
| 1851 |
+
return imat ? imat[blk * QK_K + i] : 1.0f;
|
| 1852 |
+
}
|
| 1853 |
+
|
| 1854 |
+
/* Frozen codes: put (d, dmin) on Ξ£(xβdeq) = βcarry if SSE stays in budget. */
|
| 1855 |
+
static int hex_q2k_hit_dc_carry(const float *x, const uint8_t *L,
|
| 1856 |
+
const uint8_t *scales, float *dm, float *mm,
|
| 1857 |
+
float dc_carry, const float *w256)
|
| 1858 |
+
{
|
| 1859 |
+
float d0 = *dm, m0 = *mm;
|
| 1860 |
+
double A = 0.0, B = 0.0, Sx = 0.0, sse0 = 0.0, dc0 = 0.0;
|
| 1861 |
+
for (int i = 0; i < QK_K; i++) {
|
| 1862 |
+
int j = i >> 4;
|
| 1863 |
+
float a = (float)(scales[j] & 0xF) * (float)L[i];
|
| 1864 |
+
float b = (float)(scales[j] >> 4);
|
| 1865 |
+
float wi = w256 ? w256[i] : 1.0f;
|
| 1866 |
+
float e = x[i] - (d0 * a - m0 * b);
|
| 1867 |
+
sse0 += (double)wi * e * e;
|
| 1868 |
+
dc0 += e;
|
| 1869 |
+
A += a; B += b; Sx += x[i];
|
| 1870 |
+
}
|
| 1871 |
+
double T = Sx + (double)dc_carry;
|
| 1872 |
+
float d_try = d0, m_try = m0;
|
| 1873 |
+
|
| 1874 |
+
if (fabs(B) > 1e-12) {
|
| 1875 |
+
double k = A / B, tB = T / B;
|
| 1876 |
+
double Szz = 0.0, Syz = 0.0;
|
| 1877 |
+
for (int i = 0; i < QK_K; i++) {
|
| 1878 |
+
int j = i >> 4;
|
| 1879 |
+
float a = (float)(scales[j] & 0xF) * (float)L[i];
|
| 1880 |
+
float b = (float)(scales[j] >> 4);
|
| 1881 |
+
float wi = w256 ? w256[i] : 1.0f;
|
| 1882 |
+
double z = (double)a - k * (double)b;
|
| 1883 |
+
double y = (double)x[i] - tB * (double)b;
|
| 1884 |
+
Szz += (double)wi * z * z;
|
| 1885 |
+
Syz += (double)wi * y * z;
|
| 1886 |
+
}
|
| 1887 |
+
if (Szz < 1e-30) return 0;
|
| 1888 |
+
double d_ref = Syz / Szz;
|
| 1889 |
+
double m_ref = (A * d_ref - T) / B;
|
| 1890 |
+
if (d_ref <= 0.0 || m_ref < 0.0) return 0;
|
| 1891 |
+
d_try = gguf_fp16_to_fp32(gguf_fp32_to_fp16((float)d_ref));
|
| 1892 |
+
m_try = gguf_fp16_to_fp32(gguf_fp32_to_fp16((float)m_ref));
|
| 1893 |
+
} else if (fabs(A) > 1e-12) {
|
| 1894 |
+
double d_ref = T / A;
|
| 1895 |
+
if (d_ref <= 0.0) return 0;
|
| 1896 |
+
double num = 0.0, den = 0.0;
|
| 1897 |
+
for (int i = 0; i < QK_K; i++) {
|
| 1898 |
+
int j = i >> 4;
|
| 1899 |
+
float a = (float)(scales[j] & 0xF) * (float)L[i];
|
| 1900 |
+
float b = (float)(scales[j] >> 4);
|
| 1901 |
+
float wi = w256 ? w256[i] : 1.0f;
|
| 1902 |
+
num += (double)wi * ((double)x[i] - d_ref * (double)a) * (double)b;
|
| 1903 |
+
den += (double)wi * (double)b * (double)b;
|
| 1904 |
+
}
|
| 1905 |
+
if (den < 1e-30) return 0;
|
| 1906 |
+
double m_ref = -num / den;
|
| 1907 |
+
if (m_ref < 0.0) return 0;
|
| 1908 |
+
d_try = gguf_fp16_to_fp32(gguf_fp32_to_fp16((float)d_ref));
|
| 1909 |
+
m_try = gguf_fp16_to_fp32(gguf_fp32_to_fp16((float)m_ref));
|
| 1910 |
+
} else {
|
| 1911 |
+
return 0;
|
| 1912 |
+
}
|
| 1913 |
+
if (d_try <= 0.0f || m_try < 0.0f) return 0;
|
| 1914 |
+
|
| 1915 |
+
double sse1 = 0.0, dc1 = 0.0;
|
| 1916 |
+
for (int i = 0; i < QK_K; i++) {
|
| 1917 |
+
int j = i >> 4;
|
| 1918 |
+
float a = (float)(scales[j] & 0xF) * (float)L[i];
|
| 1919 |
+
float b = (float)(scales[j] >> 4);
|
| 1920 |
+
float wi = w256 ? w256[i] : 1.0f;
|
| 1921 |
+
float e = x[i] - (d_try * a - m_try * b);
|
| 1922 |
+
sse1 += (double)wi * e * e;
|
| 1923 |
+
dc1 += e;
|
| 1924 |
+
}
|
| 1925 |
+
double off0 = dc0 + (double)dc_carry;
|
| 1926 |
+
double off1 = dc1 + (double)dc_carry;
|
| 1927 |
+
if (fabs(off1) >= fabs(off0) - 1e-12) return 0;
|
| 1928 |
+
if (sse1 > sse0 * (1.0 + (double)HEX_DC_SSE_BUDGET)) return 0;
|
| 1929 |
+
*dm = d_try;
|
| 1930 |
+
*mm = m_try;
|
| 1931 |
+
return 1;
|
| 1932 |
+
}
|
| 1933 |
+
|
| 1934 |
+
/* qΒ±1 toward Ξ£e β βcarry, spending at most HEX_DC_SSE_BUDGET extra SSE. */
|
| 1935 |
+
static void hex_q2k_dc_nudge_codes(const float *x, uint8_t *L,
|
| 1936 |
+
const uint8_t *scales, float dm, float mm,
|
| 1937 |
+
float dc_carry, const float *w256)
|
| 1938 |
+
{
|
| 1939 |
+
float e[QK_K];
|
| 1940 |
+
float sse = 0.0f, dc = 0.0f;
|
| 1941 |
+
for (int i = 0; i < QK_K; i++) {
|
| 1942 |
+
int j = i >> 4;
|
| 1943 |
+
float d_s = dm * (float)(scales[j] & 0xF);
|
| 1944 |
+
float m_s = mm * (float)(scales[j] >> 4);
|
| 1945 |
+
e[i] = x[i] - (d_s * (float)L[i] - m_s);
|
| 1946 |
+
float wi = w256 ? w256[i] : 1.0f;
|
| 1947 |
+
sse += e[i] * e[i] * wi;
|
| 1948 |
+
dc += e[i];
|
| 1949 |
+
}
|
| 1950 |
+
float cap = sse * (1.0f + HEX_DC_SSE_BUDGET);
|
| 1951 |
+
float median_step = dm * 4.0f;
|
| 1952 |
+
if (median_step < 1e-15f) median_step = 1e-15f;
|
| 1953 |
+
|
| 1954 |
+
for (int pass = 0; pass < 64; pass++) {
|
| 1955 |
+
float off = dc + dc_carry;
|
| 1956 |
+
if (fabsf(off) <= median_step) break;
|
| 1957 |
+
int best_i = -1, best_q = 0;
|
| 1958 |
+
float best_ratio = 0.0f;
|
| 1959 |
+
for (int i = 0; i < QK_K; i++) {
|
| 1960 |
+
int j = i >> 4;
|
| 1961 |
+
float d_s = dm * (float)(scales[j] & 0xF);
|
| 1962 |
+
float m_s = mm * (float)(scales[j] >> 4);
|
| 1963 |
+
if (d_s < 1e-15f) continue;
|
| 1964 |
+
int q_cur = (int)L[i];
|
| 1965 |
+
int q_try = (off > 0.0f) ? q_cur + 1 : q_cur - 1;
|
| 1966 |
+
if (q_try < 0 || q_try > 3) continue;
|
| 1967 |
+
float e_new = x[i] - (d_s * (float)q_try - m_s);
|
| 1968 |
+
float dc_red = fabsf(off) - fabsf(off + (e_new - e[i]));
|
| 1969 |
+
if (dc_red <= 0.0f) continue;
|
| 1970 |
+
float wi = w256 ? w256[i] : 1.0f;
|
| 1971 |
+
float sse_new = sse + wi * (e_new * e_new - e[i] * e[i]);
|
| 1972 |
+
if (sse_new > cap) continue;
|
| 1973 |
+
float sse_cost = sse_new - sse;
|
| 1974 |
+
if (sse_cost < 0.0f) sse_cost = 0.0f;
|
| 1975 |
+
float ratio = dc_red / (sse_cost + 1e-20f);
|
| 1976 |
+
if (ratio > best_ratio) {
|
| 1977 |
+
best_ratio = ratio;
|
| 1978 |
+
best_i = i;
|
| 1979 |
+
best_q = q_try;
|
| 1980 |
+
}
|
| 1981 |
+
}
|
| 1982 |
+
if (best_i < 0) break;
|
| 1983 |
+
{
|
| 1984 |
+
int j = best_i >> 4;
|
| 1985 |
+
float d_s = dm * (float)(scales[j] & 0xF);
|
| 1986 |
+
float m_s = mm * (float)(scales[j] >> 4);
|
| 1987 |
+
float e_new = x[best_i] - (d_s * (float)best_q - m_s);
|
| 1988 |
+
float wi = w256 ? w256[best_i] : 1.0f;
|
| 1989 |
+
sse += wi * (e_new * e_new - e[best_i] * e[best_i]);
|
| 1990 |
+
dc += (e_new - e[best_i]);
|
| 1991 |
+
e[best_i] = e_new;
|
| 1992 |
+
L[best_i] = (uint8_t)best_q;
|
| 1993 |
+
}
|
| 1994 |
+
}
|
| 1995 |
+
}
|
| 1996 |
+
|
| 1997 |
/* Robust temperature estimator for the HExState measurement model.
|
| 1998 |
*
|
| 1999 |
* The old path estimated T from the mean of each block's MAXIMUM candidate
|
|
|
|
| 3793 |
}
|
| 3794 |
|
| 3795 |
/* ββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββ
|
| 3796 |
+
* PHASE 3.9 β ROLLING DC *RESIDUAL CARRY* (does NOT shift weights)
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 3797 |
*
|
| 3798 |
+
* Bug that killed PPL: we used to quantize xβ² = x β bias so llama
|
| 3799 |
+
* stored a DC-shifted matrix. Cancellation belongs on the residual
|
| 3800 |
+
* e = x β deq of an unshifted reconstruction:
|
| 3801 |
*
|
| 3802 |
+
* carry[N] = DC_DECAY Β· Ξ£ e_{Nβ1}
|
| 3803 |
+
* prefer Ξ£ e_N β βcarry[N]
|
| 3804 |
*
|
| 3805 |
+
* via the spectral term (οΏ½οΏ½e + carry)Β², while every code/scale is
|
| 3806 |
+
* still chosen against the true x.
|
|
|
|
|
|
|
|
|
|
|
|
|
| 3807 |
* ββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββ */
|
| 3808 |
|
| 3809 |
#define DC_DECAY (g_hex_dc_decay)
|
| 3810 |
|
| 3811 |
+
float *block_dc_carry = (float *)calloc(n_blocks, sizeof(float));
|
| 3812 |
|
| 3813 |
+
if (block_dc_carry) {
|
| 3814 |
float rolling_dc = 0.0f;
|
|
|
|
|
|
|
|
|
|
|
|
|
| 3815 |
int64_t blocks_per_row = (row_width > 0 && row_width % QK_K == 0)
|
| 3816 |
? row_width / QK_K : 0;
|
| 3817 |
|
| 3818 |
for (int64_t blk = 0; blk < n_blocks; blk++) {
|
|
|
|
| 3819 |
if (blocks_per_row > 0 && (blk % blocks_per_row) == 0)
|
| 3820 |
rolling_dc = 0.0f;
|
| 3821 |
|
|
|
|
| 3830 |
hex_derive_subscales(seeds[blk].scales, seeds[blk].mins,
|
| 3831 |
dm0, mm0, dc_Ls, dc_Lm);
|
| 3832 |
|
| 3833 |
+
/* Residual-space carry: next block should cancel this, not
|
| 3834 |
+
* reconstruct a shifted x. */
|
| 3835 |
+
block_dc_carry[blk] = DC_DECAY * rolling_dc;
|
| 3836 |
|
|
|
|
|
|
|
|
|
|
| 3837 |
float dc_res = 0.0f;
|
| 3838 |
int j, k;
|
| 3839 |
for (j = 0; j < N_SUB; j++) {
|
| 3840 |
float d_sub = dm0 * (float)dc_Ls[j];
|
| 3841 |
float m_sub = mm0 * (float)dc_Lm[j];
|
| 3842 |
for (k = 0; k < 16; k++) {
|
| 3843 |
+
float x = bx[16*j + k];
|
| 3844 |
int q = 0;
|
| 3845 |
if (d_sub >= 1e-15f) {
|
| 3846 |
+
q = gguf_nearest_int((x + m_sub) / d_sub);
|
| 3847 |
if (q < 0) q = 0;
|
| 3848 |
if (q > 3) q = 3;
|
| 3849 |
}
|
| 3850 |
float deq = d_sub * (float)q - m_sub;
|
| 3851 |
+
dc_res += x - deq;
|
|
|
|
| 3852 |
}
|
| 3853 |
}
|
| 3854 |
rolling_dc = dc_res;
|
|
|
|
| 3872 |
const float *block_x = weights + blk * QK_K;
|
| 3873 |
int cidx = best_candidate[blk];
|
| 3874 |
uint8_t Ls_blk[16], Lm_blk[16];
|
| 3875 |
+
const float dc_carry = (block_dc_carry) ? block_dc_carry[blk] : 0.0f;
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 3876 |
|
| 3877 |
uint16_t base_c_d16, base_c_m16;
|
| 3878 |
hex_candidate_pair(seeds[blk].base_dm, seeds[blk].base_mm, cidx, &base_c_d16, &base_c_m16);
|
|
|
|
| 3890 |
float state_err[N_SUB][6];
|
| 3891 |
|
| 3892 |
for (int j = 0; j < N_SUB; j++) {
|
| 3893 |
+
const float *sx = block_x + 16 * j;
|
| 3894 |
for (int v = 0; v < 6; v++) state_err[j][v] = 1e30f;
|
| 3895 |
|
| 3896 |
for (int try_ls = 0; try_ls <= 15; try_ls++) {
|
|
|
|
| 3998 |
continue;
|
| 3999 |
}
|
| 4000 |
for (int k = 0; k < 16; k++) {
|
| 4001 |
+
int q = gguf_nearest_int((block_x[16*j+k] + m_sub) / d_sub);
|
| 4002 |
if (q < 0) q = 0; if (q > 3) q = 3;
|
| 4003 |
L[16*j+k] = (uint8_t)q;
|
| 4004 |
}
|
|
|
|
| 4009 |
float ls_f = (float)Ls_blk[j];
|
| 4010 |
float lm_f = (float)Lm_blk[j];
|
| 4011 |
for (int k = 0; k < 16; k++) {
|
| 4012 |
+
float x = block_x[16*j+k];
|
| 4013 |
float w = (imat_importance) ?
|
| 4014 |
imat_importance[blk * QK_K + 16*j+k] : 1.0f;
|
| 4015 |
float a = ls_f * (float)L[16*j+k];
|
|
|
|
| 4065 |
float d_sub = trial_dm * (float)Ls_blk[j];
|
| 4066 |
float m_sub = trial_mm * (float)Lm_blk[j];
|
| 4067 |
for (int k = 0; k < 16; k++) {
|
| 4068 |
+
float x = block_x[16*j+k];
|
| 4069 |
float w = (imat_importance) ?
|
| 4070 |
imat_importance[blk * QK_K + 16*j+k] : 1.0f;
|
| 4071 |
int q;
|
|
|
|
| 4091 |
}
|
| 4092 |
|
| 4093 |
for (int j = 0; j < N_SUB; j++) {
|
| 4094 |
+
const float *sx = block_x + 16 * j;
|
| 4095 |
float best_sub_err = 1e30f;
|
| 4096 |
uint8_t best_ls = Ls_blk[j], best_lm = Lm_blk[j];
|
| 4097 |
for (int try_ls = 0; try_ls <= 15; try_ls++) {
|
|
|
|
| 4158 |
q_cont_all[i] = 0.0f;
|
| 4159 |
q_base_all[i] = 0;
|
| 4160 |
} else {
|
| 4161 |
+
float qc = (block_x[i] + m_s) / d_s;
|
|
|
|
| 4162 |
q_cont_all[i] = qc;
|
| 4163 |
int qr = gguf_nearest_int(qc);
|
| 4164 |
if (qr < 0) qr = 0; if (qr > 3) qr = 3;
|
|
|
|
| 4168 |
memcpy(q_shaped_all, q_base_all, QK_K * sizeof(int));
|
| 4169 |
|
| 4170 |
float e_live[QK_K];
|
| 4171 |
+
float dc_cur = dc_carry;
|
| 4172 |
+
float sse_live = 0.0f;
|
| 4173 |
for (int i = 0; i < QK_K; i++) {
|
| 4174 |
int jj = i >> 4;
|
| 4175 |
float d_s = dm * (float)(output[blk].scales[jj] & 0xF);
|
| 4176 |
float m_s = mm * (float)(output[blk].scales[jj] >> 4);
|
| 4177 |
float deq = d_s * (float)q_shaped_all[i] - m_s;
|
|
|
|
|
|
|
| 4178 |
e_live[i] = block_x[i] - deq;
|
| 4179 |
+
dc_cur += e_live[i];
|
| 4180 |
+
float w = (imat_importance) ? imat_importance[blk * QK_K + i] : 1.0f;
|
| 4181 |
+
sse_live += e_live[i] * e_live[i] * w;
|
| 4182 |
}
|
| 4183 |
+
const float sse_cap = sse_live * (1.0f + HEX_DC_SSE_BUDGET);
|
| 4184 |
|
| 4185 |
float v_live[QK_K / 2];
|
| 4186 |
float vesica_cur = 0.0f;
|
|
|
|
| 4213 |
if (q_try == q_cur) continue;
|
| 4214 |
|
| 4215 |
float e_new = block_x[k] - (d_s * (float)q_try - m_s);
|
| 4216 |
+
float w = (imat_importance) ? imat_importance[blk * QK_K + k] : 1.0f;
|
| 4217 |
+
float sse_alt = sse_live + w * (e_new * e_new - e_live[k] * e_live[k]);
|
| 4218 |
+
if (sse_alt > sse_cap) continue;
|
| 4219 |
float de = e_new - e_live[k];
|
| 4220 |
|
| 4221 |
float v_new = v_live[pi] + de;
|
|
|
|
| 4248 |
metric_cur = 4.0f * vesica_cur + dc_cur * dc_cur;
|
| 4249 |
v_live[pi_c] = v_new_c;
|
| 4250 |
e_live[best_k]= e_new_c;
|
| 4251 |
+
{
|
| 4252 |
+
float w_c = (imat_importance) ?
|
| 4253 |
+
imat_importance[blk * QK_K + best_k] : 1.0f;
|
| 4254 |
+
sse_live += w_c * (e_new_c * e_new_c
|
| 4255 |
+
- (e_new_c - de_c) * (e_new_c - de_c));
|
| 4256 |
+
}
|
| 4257 |
}
|
| 4258 |
}
|
| 4259 |
|
| 4260 |
+
float sse_base = 0.0f, sse_shaped = 0.0f;
|
|
|
|
| 4261 |
float e_qb[QK_K], e_qs[QK_K];
|
| 4262 |
for (int i = 0; i < QK_K; i++) {
|
| 4263 |
int jj = i >> 4;
|
| 4264 |
float d_s = dm * (float)(output[blk].scales[jj] & 0xF);
|
| 4265 |
float m_s = mm * (float)(output[blk].scales[jj] >> 4);
|
| 4266 |
float w = (imat_importance) ? imat_importance[blk * QK_K + i] : 1.0f;
|
| 4267 |
+
float deq_b = d_s * (float)q_base_all[i] - m_s;
|
| 4268 |
float deq_s = d_s * (float)q_shaped_all[i] - m_s;
|
| 4269 |
+
float xv = block_x[i];
|
| 4270 |
e_qb[i] = xv - deq_b;
|
| 4271 |
e_qs[i] = xv - deq_s;
|
| 4272 |
+
sse_base += e_qb[i] * e_qb[i] * w;
|
| 4273 |
+
sse_shaped += e_qs[i] * e_qs[i] * w;
|
| 4274 |
}
|
| 4275 |
+
float err_base = sse_base + hex_spectral_penalty_ex(e_qb, QK_K, dc_carry);
|
| 4276 |
+
float err_shaped = sse_shaped + hex_spectral_penalty_ex(e_qs, QK_K, dc_carry);
|
| 4277 |
{
|
| 4278 |
+
int use_shaped = (sse_shaped <= sse_base * (1.0f + HEX_DC_SSE_BUDGET)
|
| 4279 |
+
&& err_shaped <= err_base);
|
| 4280 |
for (int i = 0; i < QK_K; i++)
|
| 4281 |
L[i] = (uint8_t)(use_shaped ? q_shaped_all[i] : q_base_all[i]);
|
| 4282 |
}
|
|
|
|
| 4315 |
for (fs_k = 0; fs_k < 16; fs_k++) {
|
| 4316 |
int idx = base + fs_k;
|
| 4317 |
float x_orig = block_x[idx];
|
| 4318 |
+
float x_adj = block_x[idx] + carry; /* adjusted + diffused */
|
| 4319 |
|
| 4320 |
/* Propose new code from diffused target */
|
| 4321 |
int q_fs = gguf_nearest_int((x_adj + m_s) / d_s);
|
|
|
|
| 4336 |
/* Propagate 7/16 of the residual (adj target vs committed code) */
|
| 4337 |
{
|
| 4338 |
float deq_final = d_s * (float)L[idx] - m_s;
|
| 4339 |
+
float residual = (block_x[idx] - deq_final);
|
| 4340 |
carry = (fs_k < 15) ? residual * (7.0f / 16.0f) : 0.0f;
|
| 4341 |
}
|
| 4342 |
}
|
|
|
|
| 4357 |
{
|
| 4358 |
/* Compute current whole-block DC and per-weight residuals */
|
| 4359 |
float wb_e[QK_K];
|
| 4360 |
+
float wb_dc = 0.0f, wb_sse = 0.0f;
|
| 4361 |
for (int i = 0; i < QK_K; i++) {
|
| 4362 |
int jj = i >> 4;
|
| 4363 |
float d_s = dm * (float)(output[blk].scales[jj] & 0xF);
|
|
|
|
| 4365 |
float deq = d_s * (float)L[i] - m_s;
|
| 4366 |
wb_e[i] = block_x[i] - deq;
|
| 4367 |
wb_dc += wb_e[i];
|
| 4368 |
+
float w = (imat_importance) ? imat_importance[blk * QK_K + i] : 1.0f;
|
| 4369 |
+
wb_sse += wb_e[i] * wb_e[i] * w;
|
| 4370 |
}
|
| 4371 |
|
| 4372 |
+
float wb_off = wb_dc + dc_carry;
|
| 4373 |
+
float median_step = dm * 4.0f;
|
|
|
|
| 4374 |
if (median_step < 1e-15f) median_step = 1e-15f;
|
| 4375 |
+
float wb_cap = wb_sse * (1.0f + HEX_DC_SSE_BUDGET);
|
| 4376 |
|
| 4377 |
for (int dc_pass = 0; dc_pass < 32; dc_pass++) {
|
| 4378 |
+
if (fabsf(wb_off) <= median_step) break;
|
| 4379 |
|
|
|
|
|
|
|
| 4380 |
int best_i = -1;
|
| 4381 |
int best_q = 0;
|
| 4382 |
+
float best_ratio = 0.0f;
|
| 4383 |
|
| 4384 |
for (int i = 0; i < QK_K; i++) {
|
| 4385 |
int jj = i >> 4;
|
|
|
|
| 4387 |
float m_s = mm * (float)(output[blk].scales[jj] >> 4);
|
| 4388 |
if (d_s < 1e-15f) continue;
|
| 4389 |
|
|
|
|
|
|
|
| 4390 |
int q_cur = (int)L[i];
|
| 4391 |
+
int q_try = (wb_off > 0.0f) ? q_cur + 1 : q_cur - 1;
|
| 4392 |
if (q_try < 0 || q_try > 3) continue;
|
| 4393 |
|
| 4394 |
float deq_new = d_s * (float)q_try - m_s;
|
| 4395 |
float e_new = block_x[i] - deq_new;
|
| 4396 |
+
float dc_reduction = fabsf(wb_off) - fabsf(wb_off + (e_new - wb_e[i]));
|
| 4397 |
if (dc_reduction <= 0.0f) continue;
|
| 4398 |
|
| 4399 |
float w = (imat_importance) ?
|
| 4400 |
imat_importance[blk * QK_K + i] : 1.0f;
|
| 4401 |
+
float sse_new = wb_sse + w * (e_new * e_new - wb_e[i] * wb_e[i]);
|
| 4402 |
+
if (sse_new > wb_cap) continue;
|
| 4403 |
+
float sse_cost = sse_new - wb_sse;
|
| 4404 |
+
if (sse_cost < 0.0f) sse_cost = 0.0f;
|
| 4405 |
+
float ratio = dc_reduction / (sse_cost + 1e-20f);
|
|
|
|
|
|
|
|
|
|
|
|
|
| 4406 |
if (ratio > best_ratio) {
|
| 4407 |
best_ratio = ratio;
|
| 4408 |
best_i = i;
|
|
|
|
| 4410 |
}
|
| 4411 |
}
|
| 4412 |
|
| 4413 |
+
if (best_i < 0) break;
|
| 4414 |
|
|
|
|
|
|
|
| 4415 |
{
|
| 4416 |
int jj = best_i >> 4;
|
| 4417 |
float d_s = dm * (float)(output[blk].scales[jj] & 0xF);
|
| 4418 |
float m_s = mm * (float)(output[blk].scales[jj] >> 4);
|
| 4419 |
float deq_new = d_s * (float)best_q - m_s;
|
| 4420 |
float e_new = block_x[best_i] - deq_new;
|
| 4421 |
+
float w = (imat_importance) ?
|
| 4422 |
+
imat_importance[blk * QK_K + best_i] : 1.0f;
|
| 4423 |
+
wb_sse += w * (e_new * e_new - wb_e[best_i] * wb_e[best_i]);
|
| 4424 |
+
wb_dc += (e_new - wb_e[best_i]);
|
| 4425 |
+
wb_off = wb_dc + dc_carry;
|
| 4426 |
wb_e[best_i] = e_new;
|
| 4427 |
L[best_i] = (uint8_t)best_q;
|
| 4428 |
}
|
|
|
|
| 4465 |
* Ξ»_dc/n; vesica/wave handled by the extended-E acceptance. */
|
| 4466 |
{
|
| 4467 |
double rw = (double)HEX_DC_LAMBDA / (double)QK_K;
|
| 4468 |
+
double rSt = rS + (double)dc_carry;
|
| 4469 |
rSaa += rw * rA * rA; rSab += rw * rA * rB;
|
| 4470 |
+
rSbb += rw * rB * rB; rSxa += rw * rSt * rA;
|
| 4471 |
+
rSxb += rw * rSt * rB;
|
| 4472 |
}
|
| 4473 |
double rdet = rSaa * rSbb - rSab * rSab;
|
| 4474 |
if (fabs(rdet) > 1e-30) {
|
|
|
|
| 4498 |
err_try += e_rt[idx] * e_rt[idx] * w;
|
| 4499 |
}
|
| 4500 |
}
|
| 4501 |
+
float sse_cur = err_cur, sse_try = err_try;
|
| 4502 |
+
err_cur += hex_spectral_penalty_ex(e_rc, QK_K, dc_carry);
|
| 4503 |
+
err_try += hex_spectral_penalty_ex(e_rt, QK_K, dc_carry);
|
| 4504 |
+
if (sse_try <= sse_cur && err_try < err_cur) {
|
| 4505 |
+
dm = dm_try; mm = mm_try;
|
| 4506 |
+
}
|
| 4507 |
}
|
| 4508 |
}
|
| 4509 |
output[blk].d = gguf_fp32_to_fp16(dm);
|
|
|
|
| 4515 |
* Objective-function mismatch fix: the final passes that commit the
|
| 4516 |
* 2-bit codes β the 16Γ16 (ls, lm) sub-block search, the Β±8 ULP
|
| 4517 |
* (d, dmin) neighborhood search, and the greedy-descent error shaping
|
| 4518 |
+
* β all minimise error against the DC-ADJUSTED target block_x.
|
| 4519 |
* The reported RMSE, however, is measured against the ORIGINAL
|
| 4520 |
* weights. The codes are therefore stranded at the optimum of a
|
| 4521 |
* SHIFTED objective, while only the scalar (d, dmin) refit above
|
|
|
|
| 4601 |
|
| 4602 |
/* Extended score of the CURRENT committed state */
|
| 4603 |
float best_sub = sub_sse[j]
|
| 4604 |
+
+ (HEX_DC_LAMBDA / (float)QK_K)
|
| 4605 |
+
* (dc_tot + dc_carry) * (dc_tot + dc_carry)
|
| 4606 |
+ (HEX_VW_LAMBDA / (float)QK_K) * ves_tot;
|
| 4607 |
int best_ls = -1, best_lm = 0;
|
| 4608 |
uint8_t best_q[16];
|
|
|
|
| 4639 |
if (sub_err >= best_sub) { aborted = 1; break; }
|
| 4640 |
}
|
| 4641 |
if (aborted) continue;
|
| 4642 |
+
if (sub_err > sub_sse[j]) continue;
|
| 4643 |
+
float dcc_tot = dc_rest + dcc + dc_carry;
|
| 4644 |
float score = sub_err
|
| 4645 |
+ (HEX_DC_LAMBDA / (float)QK_K)
|
| 4646 |
+
* dcc_tot * dcc_tot
|
| 4647 |
+ (HEX_VW_LAMBDA / (float)QK_K)
|
| 4648 |
* (ves_rest + vesc);
|
| 4649 |
if (score < best_sub) {
|
|
|
|
| 4697 |
}
|
| 4698 |
{
|
| 4699 |
double pw = (double)HEX_DC_LAMBDA / (double)QK_K;
|
| 4700 |
+
double pSt = pS + (double)dc_carry;
|
| 4701 |
pSaa += pw * pA * pA; pSab += pw * pA * pB;
|
| 4702 |
+
pSbb += pw * pB * pB; pSxa += pw * pSt * pA;
|
| 4703 |
+
pSxb += pw * pSt * pB;
|
| 4704 |
}
|
| 4705 |
double pdet = pSaa * pSbb - pSab * pSab;
|
| 4706 |
if (fabs(pdet) > 1e-30) {
|
|
|
|
| 4732 |
err_try += e_pt[idx] * e_pt[idx] * w;
|
| 4733 |
}
|
| 4734 |
}
|
| 4735 |
+
float sse_cur = err_cur, sse_try = err_try;
|
| 4736 |
+
err_cur += hex_spectral_penalty_ex(e_pc, QK_K, dc_carry);
|
| 4737 |
+
err_try += hex_spectral_penalty_ex(e_pt, QK_K, dc_carry);
|
| 4738 |
+
if (sse_try <= sse_cur && err_try < err_cur) {
|
| 4739 |
dm = dm_try;
|
| 4740 |
mm = mm_try;
|
| 4741 |
pol_improved = 1;
|
|
|
|
| 4772 |
cur_err += e_u[idx] * e_u[idx] * w;
|
| 4773 |
}
|
| 4774 |
}
|
| 4775 |
+
cur_err += hex_spectral_penalty_ex(e_u, QK_K, dc_carry);
|
| 4776 |
|
| 4777 |
float best_err = cur_err;
|
| 4778 |
uint16_t best_d16 = base_d16, best_m16 = base_m16;
|
|
|
|
| 4802 |
if (err >= best_err) { pruned = 1; break; }
|
| 4803 |
}
|
| 4804 |
if (pruned) continue;
|
| 4805 |
+
err += hex_spectral_penalty_ex(e_u, QK_K, dc_carry);
|
| 4806 |
if (err < best_err) {
|
| 4807 |
best_err = err;
|
| 4808 |
best_d16 = (uint16_t)cd16;
|
|
|
|
| 4853 |
fin_err += e_f[idx] * e_f[idx] * w;
|
| 4854 |
}
|
| 4855 |
}
|
| 4856 |
+
/* Floor is reconstruction SSE only. Spectral must not replace a
|
| 4857 |
+
* better W with a worse candidate just to zero DC. */
|
| 4858 |
float g_best = candidate_errors[blk][0];
|
| 4859 |
int g_cand = 0;
|
| 4860 |
for (int c = 1; c < TOTAL_SCALE_CANDIDATES; c++) {
|
|
|
|
| 4911 |
total_err += berr;
|
| 4912 |
}
|
| 4913 |
|
| 4914 |
+
/* ββ PHASE 4.8: sequential TRUE residual carry ββββββββββββββββββββββ
|
| 4915 |
+
* Phase 3.9 carry is a nearest-round estimate on the seed (d,dmin);
|
| 4916 |
+
* Phase 4 then rewrites codes in parallel, so that carry is stale.
|
| 4917 |
+
* Walk each row in order, measure the encoded Ξ£e, and spend a tiny
|
| 4918 |
+
* SSE budget on (d,dmin) + qΒ±1 to hit Ξ£e β βdecayΒ·R_prev. */
|
| 4919 |
+
if (HEX_DC_LAMBDA > 0.0f || DC_DECAY > 0.0f) {
|
| 4920 |
+
int64_t bpr = (row_width > 0 && row_width % QK_K == 0)
|
| 4921 |
+
? row_width / QK_K : 0;
|
| 4922 |
+
float rolling_dc = 0.0f;
|
| 4923 |
+
uint8_t L8[QK_K];
|
| 4924 |
+
float w256[QK_K];
|
| 4925 |
+
for (int64_t blk = 0; blk < n_blocks; blk++) {
|
| 4926 |
+
if (bpr > 0 && (blk % bpr) == 0)
|
| 4927 |
+
rolling_dc = 0.0f;
|
| 4928 |
+
const float *bx = weights + blk * QK_K;
|
| 4929 |
+
float carry = DC_DECAY * rolling_dc;
|
| 4930 |
+
hex_q2k_unpack_L(&output[blk], L8);
|
| 4931 |
+
float dm8 = gguf_fp16_to_fp32(output[blk].d);
|
| 4932 |
+
float mm8 = gguf_fp16_to_fp32(output[blk].dmin);
|
| 4933 |
+
for (int i = 0; i < QK_K; i++)
|
| 4934 |
+
w256[i] = hex_q2k_el_w(imat_importance, blk, i);
|
| 4935 |
+
if (hex_q2k_hit_dc_carry(bx, L8, output[blk].scales, &dm8, &mm8,
|
| 4936 |
+
carry, w256)) {
|
| 4937 |
+
output[blk].d = gguf_fp32_to_fp16(dm8);
|
| 4938 |
+
output[blk].dmin = gguf_fp32_to_fp16(mm8);
|
| 4939 |
+
}
|
| 4940 |
+
hex_q2k_dc_nudge_codes(bx, L8, output[blk].scales, dm8, mm8,
|
| 4941 |
+
carry, w256);
|
| 4942 |
+
hex_q2k_pack_L(&output[blk], L8);
|
| 4943 |
+
|
| 4944 |
+
float deq[QK_K];
|
| 4945 |
+
gguf_dequantize_q2_k_block(&output[blk], deq);
|
| 4946 |
+
float dc_res = 0.0f;
|
| 4947 |
+
for (int i = 0; i < QK_K; i++)
|
| 4948 |
+
dc_res += bx[i] - deq[i];
|
| 4949 |
+
rolling_dc = dc_res;
|
| 4950 |
+
}
|
| 4951 |
+
total_err = 0.0f;
|
| 4952 |
+
for (int64_t blk = 0; blk < n_blocks; blk++)
|
| 4953 |
+
total_err += gguf_q2_k_block_error(weights + blk * QK_K, &output[blk]);
|
| 4954 |
+
}
|
| 4955 |
+
|
| 4956 |
for (int _ti = 0; _ti < _n_omp_threads; _ti++)
|
| 4957 |
hpc_destroy(_tl_graphs[_ti]);
|
| 4958 |
free(_tl_graphs);
|
| 4959 |
|
| 4960 |
+
free(block_dc_carry);
|
| 4961 |
free(seeds);
|
| 4962 |
free(candidate_errors);
|
| 4963 |
free(best_candidate);
|
|
|
|
| 5541 |
return -1;
|
| 5542 |
}
|
| 5543 |
|
| 5544 |
+
/* βββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββ
|
| 5545 |
+
* IQ2_XS β E8-CODEBOOK QUANTIZER WITH FOLD/DC SHAPING (2.3125 bpw)
|
| 5546 |
+
*
|
| 5547 |
+
* Each 8-weight group is one of 512 codewords (magnitudes {8,25,43}) times
|
| 5548 |
+
* a sign pattern with EVEN parity (7 bits stored, 8th = parity). Sub-block
|
| 5549 |
+
* (16) scale db = dΒ·(ls+0.5)/4, ls β 0..15, d fp16 per 256-block.
|
| 5550 |
+
*
|
| 5551 |
+
* Why the codebook is where fold finally pays: for a given group there are
|
| 5552 |
+
* several codewords within a hair of the nearest one (E8 shells are dense),
|
| 5553 |
+
* so residual shaping β Ξ£e β βcarry across blocks, e_i + e_{i+128} small β
|
| 5554 |
+
* can pick among them at ~zero SSE cost. Q2_K only had Β±1 scalar steps.
|
| 5555 |
+
*
|
| 5556 |
+
* Over ggml's encoder: exact grid magnitudes in the objective, sign-parity
|
| 5557 |
+
* flip chosen jointly with the codeword, codewords re-picked after the
|
| 5558 |
+
* 4-bit scale quantisation, d candidate search, ls Β±1 descent, and the
|
| 5559 |
+
* rolling residual carry. Reconstruction target is always the true x.
|
| 5560 |
+
* βββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββ */
|
| 5561 |
+
|
| 5562 |
+
#include "iq2xs_grid.h"
|
| 5563 |
+
|
| 5564 |
+
#define IQ2XS_NGRID 512
|
| 5565 |
+
#define IQ2XS_TOPK 8
|
| 5566 |
+
#define IQ2XS_NGROUP (QK_K / 8) /* 32 */
|
| 5567 |
+
#define IQ2XS_NSUB (QK_K / 16) /* 16 */
|
| 5568 |
+
|
| 5569 |
+
static float g_iq2xs_gridf[IQ2XS_NGRID][8];
|
| 5570 |
+
static int g_iq2xs_grid_ready = 0;
|
| 5571 |
+
|
| 5572 |
+
static void iq2xs_prepare_grid(void)
|
| 5573 |
+
{
|
| 5574 |
+
if (g_iq2xs_grid_ready) return;
|
| 5575 |
+
for (int k = 0; k < IQ2XS_NGRID; k++)
|
| 5576 |
+
for (int j = 0; j < 8; j++)
|
| 5577 |
+
g_iq2xs_gridf[k][j] = (float)((iq2xs_grid[k] >> (8 * j)) & 0xFF);
|
| 5578 |
+
g_iq2xs_grid_ready = 1;
|
| 5579 |
+
}
|
| 5580 |
+
|
| 5581 |
+
typedef struct {
|
| 5582 |
+
uint16_t code; /* grid | signs7 << 9 */
|
| 5583 |
+
float sse; /* Ξ£ w (x β deq)Β² */
|
| 5584 |
+
float deq[8];
|
| 5585 |
+
} IQ2Cand;
|
| 5586 |
+
|
| 5587 |
+
/* Top-K codewords for 8 weights at magnitude scale db (db > 0).
|
| 5588 |
+
* Sign parity is enforced by the cheapest flip *per codeword*. */
|
| 5589 |
+
static int iq2xs_group_candidates(const float *x, const float *w, float db,
|
| 5590 |
+
int K, IQ2Cand *out)
|
| 5591 |
+
{
|
| 5592 |
+
float ax[8], wa[8];
|
| 5593 |
+
uint8_t s = 0; int par = 0;
|
| 5594 |
+
for (int i = 0; i < 8; i++) {
|
| 5595 |
+
ax[i] = fabsf(x[i]);
|
| 5596 |
+
wa[i] = w[i];
|
| 5597 |
+
if (x[i] < 0.0f) { s |= (uint8_t)(1u << i); par ^= 1; }
|
| 5598 |
+
}
|
| 5599 |
+
int n = 0;
|
| 5600 |
+
for (int k = 0; k < IQ2XS_NGRID; k++) {
|
| 5601 |
+
const float *g = g_iq2xs_gridf[k];
|
| 5602 |
+
float e0 = 0.0f;
|
| 5603 |
+
for (int i = 0; i < 8; i++) {
|
| 5604 |
+
float d = ax[i] - db * g[i];
|
| 5605 |
+
e0 += wa[i] * d * d;
|
| 5606 |
+
}
|
| 5607 |
+
/* Parity: odd sign count needs one flip. Offer the two cheapest
|
| 5608 |
+
* flips as separate candidates β flipping a different weight is
|
| 5609 |
+
* the cheapest DC lever this offset-free format has. */
|
| 5610 |
+
int flips[2] = { -1, -1 };
|
| 5611 |
+
float costs[2] = { 0.0f, 0.0f };
|
| 5612 |
+
int nf = 1;
|
| 5613 |
+
if (par) {
|
| 5614 |
+
float c1 = 1e30f, c2 = 1e30f; int f1 = -1, f2 = -1;
|
| 5615 |
+
for (int i = 0; i < 8; i++) {
|
| 5616 |
+
float c = 4.0f * wa[i] * ax[i] * db * g[i];
|
| 5617 |
+
if (c < c1) { c2 = c1; f2 = f1; c1 = c; f1 = i; }
|
| 5618 |
+
else if (c < c2) { c2 = c; f2 = i; }
|
| 5619 |
+
}
|
| 5620 |
+
flips[0] = f1; costs[0] = c1;
|
| 5621 |
+
flips[1] = f2; costs[1] = c2;
|
| 5622 |
+
nf = (K > 1 && f2 >= 0) ? 2 : 1;
|
| 5623 |
+
}
|
| 5624 |
+
for (int f = 0; f < nf; f++) {
|
| 5625 |
+
float e = e0 + costs[f];
|
| 5626 |
+
if (n == K && e >= out[K - 1].sse) continue;
|
| 5627 |
+
int pos = n < K ? n : K - 1;
|
| 5628 |
+
while (pos > 0 && out[pos - 1].sse > e) { out[pos] = out[pos - 1]; pos--; }
|
| 5629 |
+
uint8_t sf = s;
|
| 5630 |
+
if (flips[f] >= 0) sf ^= (uint8_t)(1u << flips[f]);
|
| 5631 |
+
out[pos].code = (uint16_t)(k | ((sf & 127) << 9));
|
| 5632 |
+
out[pos].sse = e;
|
| 5633 |
+
for (int i = 0; i < 8; i++)
|
| 5634 |
+
out[pos].deq[i] = db * g[i] * ((sf >> i) & 1 ? -1.0f : 1.0f);
|
| 5635 |
+
if (n < K) n++;
|
| 5636 |
+
}
|
| 5637 |
+
}
|
| 5638 |
+
return n;
|
| 5639 |
+
}
|
| 5640 |
+
|
| 5641 |
+
/* Decode one group from its 16-bit code at scale db. */
|
| 5642 |
+
static inline void iq2xs_decode_group(uint16_t code, float db, float *deq)
|
| 5643 |
+
{
|
| 5644 |
+
const float *g = g_iq2xs_gridf[code & 511];
|
| 5645 |
+
uint8_t signs = ksigns_iq2xs[code >> 9];
|
| 5646 |
+
for (int j = 0; j < 8; j++)
|
| 5647 |
+
deq[j] = db * g[j] * ((signs >> j) & 1 ? -1.0f : 1.0f);
|
| 5648 |
+
}
|
| 5649 |
+
|
| 5650 |
+
/* Best codes for a 16-weight sub-block at fixed db; returns weighted SSE. */
|
| 5651 |
+
static float iq2xs_sub_pick(const float *x, const float *w, float db,
|
| 5652 |
+
uint16_t code[2], float deq[16])
|
| 5653 |
+
{
|
| 5654 |
+
if (db <= 0.0f) {
|
| 5655 |
+
code[0] = code[1] = 0;
|
| 5656 |
+
float e = 0.0f;
|
| 5657 |
+
for (int i = 0; i < 16; i++) { deq[i] = 0.0f; e += w[i] * x[i] * x[i]; }
|
| 5658 |
+
return e;
|
| 5659 |
+
}
|
| 5660 |
+
IQ2Cand c;
|
| 5661 |
+
float e = 0.0f;
|
| 5662 |
+
for (int k = 0; k < 2; k++) {
|
| 5663 |
+
iq2xs_group_candidates(x + 8 * k, w + 8 * k, db, 1, &c);
|
| 5664 |
+
code[k] = c.code;
|
| 5665 |
+
memcpy(deq + 8 * k, c.deq, sizeof(c.deq));
|
| 5666 |
+
e += c.sse;
|
| 5667 |
+
}
|
| 5668 |
+
return e;
|
| 5669 |
+
}
|
| 5670 |
+
|
| 5671 |
+
/* Float scale search for one sub-block: candidate db grid + LS refit. */
|
| 5672 |
+
static float iq2xs_sub_fit(const float *x, const float *w, float *db_out)
|
| 5673 |
+
{
|
| 5674 |
+
float amax = 0.0f;
|
| 5675 |
+
for (int i = 0; i < 16; i++) amax = fmaxf(amax, fabsf(x[i]));
|
| 5676 |
+
if (amax < 1e-12f) { *db_out = 0.0f; return 0.0f; }
|
| 5677 |
+
|
| 5678 |
+
float best_e = 1e30f, best_db = amax / 43.0f;
|
| 5679 |
+
uint16_t code[2]; float deq[16];
|
| 5680 |
+
for (int is = -10; is <= 10; is++) {
|
| 5681 |
+
/* amax lands on grid value 43Β·(1+0.035Β·is): includes clipped maxima */
|
| 5682 |
+
float db = amax / (43.0f * (1.0f + 0.035f * (float)is));
|
| 5683 |
+
float e = iq2xs_sub_pick(x, w, db, code, deq);
|
| 5684 |
+
/* LS refit of db with codes fixed: deq = dbΒ·Δ */
|
| 5685 |
+
double num = 0.0, den = 0.0;
|
| 5686 |
+
for (int i = 0; i < 16; i++) {
|
| 5687 |
+
double gh = deq[i] / db;
|
| 5688 |
+
num += (double)w[i] * x[i] * gh;
|
| 5689 |
+
den += (double)w[i] * gh * gh;
|
| 5690 |
+
}
|
| 5691 |
+
if (den > 0.0 && num > 0.0) {
|
| 5692 |
+
float db2 = (float)(num / den);
|
| 5693 |
+
float e2 = iq2xs_sub_pick(x, w, db2, code, deq);
|
| 5694 |
+
if (e2 < e) { e = e2; db = db2; }
|
| 5695 |
+
}
|
| 5696 |
+
if (e < best_e) { best_e = e; best_db = db; }
|
| 5697 |
+
}
|
| 5698 |
+
*db_out = best_db;
|
| 5699 |
+
return best_e;
|
| 5700 |
+
}
|
| 5701 |
+
|
| 5702 |
+
/* Whole-block encode at a given d: ls from float sub-scales, re-pick codes.
|
| 5703 |
+
* Returns weighted SSE. */
|
| 5704 |
+
static float iq2xs_block_at_d(const float *x, const float *w, float d,
|
| 5705 |
+
const float *db_f, uint8_t ls[IQ2XS_NSUB],
|
| 5706 |
+
uint16_t code[IQ2XS_NGROUP], float deq[QK_K])
|
| 5707 |
+
{
|
| 5708 |
+
float e = 0.0f;
|
| 5709 |
+
for (int ib = 0; ib < IQ2XS_NSUB; ib++) {
|
| 5710 |
+
int l = (d > 0.0f) ? gguf_nearest_int(db_f[ib] * 4.0f / d - 0.5f) : 0;
|
| 5711 |
+
if (l < 0) l = 0; if (l > 15) l = 15;
|
| 5712 |
+
ls[ib] = (uint8_t)l;
|
| 5713 |
+
float db = d * ((float)l + 0.5f) * 0.25f;
|
| 5714 |
+
e += iq2xs_sub_pick(x + 16 * ib, w + 16 * ib, db, code + 2 * ib, deq + 16 * ib);
|
| 5715 |
+
}
|
| 5716 |
+
return e;
|
| 5717 |
+
}
|
| 5718 |
+
|
| 5719 |
+
static inline float iq2xs_sub_db(float d, uint8_t ls)
|
| 5720 |
+
{
|
| 5721 |
+
return d * ((float)ls + 0.5f) * 0.25f;
|
| 5722 |
+
}
|
| 5723 |
+
|
| 5724 |
+
/* Greedy re-selection among top-K codewords per group on
|
| 5725 |
+
* SSE + (Ξ»_dc/n)(Ξ£e + carry)Β² + (Ξ»_vw/n) Ξ£_p (e_p + e_{p+128})Β²
|
| 5726 |
+
* with a block SSE cap. This is the fold-through-codebook step. */
|
| 5727 |
+
static void iq2xs_shape_block(const float *x, const float *w, float d,
|
| 5728 |
+
const uint8_t ls[IQ2XS_NSUB],
|
| 5729 |
+
uint16_t code[IQ2XS_NGROUP], float dc_carry)
|
| 5730 |
+
{
|
| 5731 |
+
if (HEX_DC_LAMBDA == 0.0f && HEX_VW_LAMBDA == 0.0f) return;
|
| 5732 |
+
|
| 5733 |
+
IQ2Cand cands[IQ2XS_NGROUP][IQ2XS_TOPK];
|
| 5734 |
+
int ncand[IQ2XS_NGROUP], cur[IQ2XS_NGROUP];
|
| 5735 |
+
float e[QK_K];
|
| 5736 |
+
float sse = 0.0f, dc = dc_carry;
|
| 5737 |
+
|
| 5738 |
+
for (int g = 0; g < IQ2XS_NGROUP; g++) {
|
| 5739 |
+
float db = iq2xs_sub_db(d, ls[g >> 1]);
|
| 5740 |
+
if (db <= 0.0f) { ncand[g] = 0; cur[g] = -1;
|
| 5741 |
+
for (int j = 0; j < 8; j++) { e[8*g+j] = x[8*g+j]; sse += w[8*g+j]*x[8*g+j]*x[8*g+j]; dc += e[8*g+j]; }
|
| 5742 |
+
continue; }
|
| 5743 |
+
ncand[g] = iq2xs_group_candidates(x + 8*g, w + 8*g, db, IQ2XS_TOPK, cands[g]);
|
| 5744 |
+
cur[g] = 0;
|
| 5745 |
+
for (int c = 0; c < ncand[g]; c++)
|
| 5746 |
+
if (cands[g][c].code == code[g]) { cur[g] = c; break; }
|
| 5747 |
+
const IQ2Cand *cc = &cands[g][cur[g]];
|
| 5748 |
+
for (int j = 0; j < 8; j++) { e[8*g+j] = x[8*g+j] - cc->deq[j]; dc += e[8*g+j]; }
|
| 5749 |
+
sse += cc->sse;
|
| 5750 |
+
}
|
| 5751 |
+
const float cap = sse * (1.0f + HEX_DC_SSE_BUDGET);
|
| 5752 |
+
float v[QK_K / 2], ves = 0.0f;
|
| 5753 |
+
for (int p = 0; p < QK_K / 2; p++) { v[p] = e[p] + e[p + QK_K/2]; ves += v[p] * v[p]; }
|
| 5754 |
+
const float ldc = HEX_DC_LAMBDA / (float)QK_K, lvw = HEX_VW_LAMBDA / (float)QK_K;
|
| 5755 |
+
float metric = sse + ldc * dc * dc + lvw * ves;
|
| 5756 |
+
|
| 5757 |
+
for (int pass = 0; pass < 96; pass++) {
|
| 5758 |
+
int best_g = -1, best_c = 0; float best_m = metric;
|
| 5759 |
+
for (int g = 0; g < IQ2XS_NGROUP; g++) {
|
| 5760 |
+
if (cur[g] < 0) continue;
|
| 5761 |
+
const IQ2Cand *co = &cands[g][cur[g]];
|
| 5762 |
+
for (int c = 0; c < ncand[g]; c++) {
|
| 5763 |
+
if (c == cur[g]) continue;
|
| 5764 |
+
const IQ2Cand *cn = &cands[g][c];
|
| 5765 |
+
float sse2 = sse - co->sse + cn->sse;
|
| 5766 |
+
if (sse2 > cap) continue;
|
| 5767 |
+
float dc2 = dc, ves2 = ves;
|
| 5768 |
+
for (int j = 0; j < 8; j++) {
|
| 5769 |
+
int i = 8*g + j;
|
| 5770 |
+
float de = co->deq[j] - cn->deq[j]; /* e_new β e_old */
|
| 5771 |
+
dc2 += de;
|
| 5772 |
+
int p = (i < QK_K/2) ? i : i - QK_K/2;
|
| 5773 |
+
float vn = v[p] + de;
|
| 5774 |
+
ves2 += vn * vn - v[p] * v[p];
|
| 5775 |
+
}
|
| 5776 |
+
float m2 = sse2 + ldc * dc2 * dc2 + lvw * ves2;
|
| 5777 |
+
if (m2 < best_m) { best_m = m2; best_g = g; best_c = c; }
|
| 5778 |
+
}
|
| 5779 |
+
}
|
| 5780 |
+
if (best_g < 0) break;
|
| 5781 |
+
const IQ2Cand *co = &cands[best_g][cur[best_g]];
|
| 5782 |
+
const IQ2Cand *cn = &cands[best_g][best_c];
|
| 5783 |
+
sse += cn->sse - co->sse;
|
| 5784 |
+
for (int j = 0; j < 8; j++) {
|
| 5785 |
+
int i = 8*best_g + j;
|
| 5786 |
+
float de = co->deq[j] - cn->deq[j];
|
| 5787 |
+
e[i] += de; dc += de;
|
| 5788 |
+
int p = (i < QK_K/2) ? i : i - QK_K/2;
|
| 5789 |
+
ves += (v[p] + de) * (v[p] + de) - v[p] * v[p];
|
| 5790 |
+
v[p] += de;
|
| 5791 |
+
}
|
| 5792 |
+
cur[best_g] = best_c;
|
| 5793 |
+
code[best_g] = cn->code;
|
| 5794 |
+
metric = best_m;
|
| 5795 |
+
}
|
| 5796 |
+
}
|
| 5797 |
+
|
| 5798 |
+
static void iq2xs_pack(BlockIQ2XS *b, float d, const uint8_t ls[IQ2XS_NSUB],
|
| 5799 |
+
const uint16_t code[IQ2XS_NGROUP])
|
| 5800 |
+
{
|
| 5801 |
+
b->d = gguf_fp32_to_fp16(d);
|
| 5802 |
+
for (int ib = 0; ib < IQ2XS_NSUB; ib += 2)
|
| 5803 |
+
b->scales[ib / 2] = (uint8_t)(ls[ib] | (ls[ib + 1] << 4));
|
| 5804 |
+
memcpy(b->qs, code, sizeof(uint16_t) * IQ2XS_NGROUP);
|
| 5805 |
+
}
|
| 5806 |
+
|
| 5807 |
+
static void iq2xs_dequant_block(const BlockIQ2XS *b, float *out)
|
| 5808 |
+
{
|
| 5809 |
+
iq2xs_prepare_grid();
|
| 5810 |
+
float d = gguf_fp16_to_fp32(b->d);
|
| 5811 |
+
for (int g = 0; g < IQ2XS_NGROUP; g++) {
|
| 5812 |
+
uint8_t ls = (g & 2) ? (b->scales[g >> 2] >> 4) : (b->scales[g >> 2] & 0xF);
|
| 5813 |
+
iq2xs_decode_group(b->qs[g], iq2xs_sub_db(d, ls), out + 8 * g);
|
| 5814 |
+
}
|
| 5815 |
+
}
|
| 5816 |
+
|
| 5817 |
+
static void quantize_tensor_iq2_xs_hpc(const float *weights, int64_t n_elements,
|
| 5818 |
+
BlockIQ2XS *output, float *out_total_error,
|
| 5819 |
+
const float *imat_importance, int verbose,
|
| 5820 |
+
int64_t row_width)
|
| 5821 |
+
{
|
| 5822 |
+
if (!weights || !output || n_elements <= 0 || n_elements % QK_K != 0) {
|
| 5823 |
+
if (out_total_error) *out_total_error = -1.0f;
|
| 5824 |
+
return;
|
| 5825 |
+
}
|
| 5826 |
+
iq2xs_prepare_grid();
|
| 5827 |
+
const int64_t n_blocks = n_elements / QK_K;
|
| 5828 |
+
static const float d_mult[] = { 1.0f, 0.97f, 1.03f, 0.94f, 1.06f, 0.90f, 1.10f, 0.85f, 1.15f };
|
| 5829 |
+
const int n_dm = (int)(sizeof(d_mult) / sizeof(d_mult[0]));
|
| 5830 |
+
|
| 5831 |
+
/* Experimental knobs, both OFF by default β measured on a controlled
|
| 5832 |
+
* splice A/B (SmolLM2 ffn_down, 64Γ512-token PPL, imatrix = E[aΒ²]):
|
| 5833 |
+
* HEX_IQ2_WMODE=1 ggml's w = imatΒ·sqrt(ΟΒ²_blk + xΒ²) β PPL +6% (worse)
|
| 5834 |
+
* HEX_IQ2_INFLATE=Ξ± scale d by (1+Ξ±) after the fit β PPL +1..4% (worse)
|
| 5835 |
+
* Plain imatrix-weighted SSE with the exact grid is the best objective. */
|
| 5836 |
+
const char *wm = getenv("HEX_IQ2_WMODE");
|
| 5837 |
+
const int wmode = wm ? atoi(wm) : 0;
|
| 5838 |
+
const char *inf = getenv("HEX_IQ2_INFLATE");
|
| 5839 |
+
const float inflate = inf ? (float)atof(inf) : 0.0f;
|
| 5840 |
+
|
| 5841 |
+
#pragma omp parallel for schedule(dynamic, 16)
|
| 5842 |
+
for (int64_t blk = 0; blk < n_blocks; blk++) {
|
| 5843 |
+
const float *x = weights + blk * QK_K;
|
| 5844 |
+
float w[QK_K];
|
| 5845 |
+
float sigma2 = 0.0f;
|
| 5846 |
+
for (int i = 0; i < QK_K; i++) sigma2 += x[i] * x[i];
|
| 5847 |
+
sigma2 /= (float)QK_K;
|
| 5848 |
+
for (int i = 0; i < QK_K; i++) {
|
| 5849 |
+
float base = imat_importance ? imat_importance[blk * QK_K + i] : 1.0f;
|
| 5850 |
+
w[i] = (wmode == 1) ? base * sqrtf(sigma2 + x[i] * x[i]) : base;
|
| 5851 |
+
}
|
| 5852 |
+
|
| 5853 |
+
/* 1. float sub-block scales */
|
| 5854 |
+
float db_f[IQ2XS_NSUB], db_max = 0.0f;
|
| 5855 |
+
for (int ib = 0; ib < IQ2XS_NSUB; ib++) {
|
| 5856 |
+
iq2xs_sub_fit(x + 16 * ib, w + 16 * ib, &db_f[ib]);
|
| 5857 |
+
db_max = fmaxf(db_max, db_f[ib]);
|
| 5858 |
+
}
|
| 5859 |
+
if (db_max <= 0.0f) { memset(&output[blk], 0, sizeof(BlockIQ2XS)); continue; }
|
| 5860 |
+
|
| 5861 |
+
/* 2. d candidate search (fp16-exact), codes re-picked at quantised db */
|
| 5862 |
+
float d0 = db_max * 4.0f / 15.5f;
|
| 5863 |
+
float best_e = 1e30f, best_d = d0;
|
| 5864 |
+
uint8_t ls[IQ2XS_NSUB], ls_t[IQ2XS_NSUB];
|
| 5865 |
+
uint16_t code[IQ2XS_NGROUP], code_t[IQ2XS_NGROUP];
|
| 5866 |
+
float deq[QK_K];
|
| 5867 |
+
for (int c = 0; c < n_dm; c++) {
|
| 5868 |
+
float d = gguf_fp16_to_fp32(gguf_fp32_to_fp16(d0 * d_mult[c]));
|
| 5869 |
+
if (d <= 0.0f) continue;
|
| 5870 |
+
float e = iq2xs_block_at_d(x, w, d, db_f, ls_t, code_t, deq);
|
| 5871 |
+
if (e < best_e) { best_e = e; best_d = d;
|
| 5872 |
+
memcpy(ls, ls_t, sizeof(ls)); memcpy(code, code_t, sizeof(code)); }
|
| 5873 |
+
}
|
| 5874 |
+
|
| 5875 |
+
/* 3. per-sub-block ls Β±1 coordinate descent at fixed d */
|
| 5876 |
+
float sub_e[IQ2XS_NSUB];
|
| 5877 |
+
for (int ib = 0; ib < IQ2XS_NSUB; ib++)
|
| 5878 |
+
sub_e[ib] = iq2xs_sub_pick(x + 16*ib, w + 16*ib, iq2xs_sub_db(best_d, ls[ib]),
|
| 5879 |
+
code + 2*ib, deq + 16*ib);
|
| 5880 |
+
for (int it = 0; it < 3; it++) {
|
| 5881 |
+
int moved = 0;
|
| 5882 |
+
for (int ib = 0; ib < IQ2XS_NSUB; ib++) {
|
| 5883 |
+
for (int dl = -1; dl <= 1; dl += 2) {
|
| 5884 |
+
int l = (int)ls[ib] + dl;
|
| 5885 |
+
if (l < 0 || l > 15) continue;
|
| 5886 |
+
uint16_t ct[2]; float dq[16];
|
| 5887 |
+
float e = iq2xs_sub_pick(x + 16*ib, w + 16*ib,
|
| 5888 |
+
iq2xs_sub_db(best_d, (uint8_t)l), ct, dq);
|
| 5889 |
+
if (e < sub_e[ib]) {
|
| 5890 |
+
sub_e[ib] = e; ls[ib] = (uint8_t)l;
|
| 5891 |
+
code[2*ib] = ct[0]; code[2*ib+1] = ct[1];
|
| 5892 |
+
memcpy(deq + 16*ib, dq, sizeof(dq));
|
| 5893 |
+
moved = 1;
|
| 5894 |
+
}
|
| 5895 |
+
}
|
| 5896 |
+
}
|
| 5897 |
+
if (!moved) break;
|
| 5898 |
+
}
|
| 5899 |
+
|
| 5900 |
+
/* 4. fold/DC shaping among near-equivalent codewords (carry = 0 here;
|
| 5901 |
+
* the sequential pass below applies the true residual carry). */
|
| 5902 |
+
iq2xs_shape_block(x, w, best_d, ls, code, 0.0f);
|
| 5903 |
+
|
| 5904 |
+
float d_out = best_d;
|
| 5905 |
+
if (inflate != 0.0f)
|
| 5906 |
+
d_out = gguf_fp16_to_fp32(gguf_fp32_to_fp16(best_d * (1.0f + inflate)));
|
| 5907 |
+
iq2xs_pack(&output[blk], d_out, ls, code);
|
| 5908 |
+
}
|
| 5909 |
+
|
| 5910 |
+
/* 5. sequential TRUE residual carry along each row */
|
| 5911 |
+
if (HEX_DC_LAMBDA > 0.0f && g_hex_dc_decay > 0.0f) {
|
| 5912 |
+
int64_t bpr = (row_width > 0 && row_width % QK_K == 0) ? row_width / QK_K : 0;
|
| 5913 |
+
float rolling = 0.0f, w[QK_K], deq[QK_K];
|
| 5914 |
+
uint8_t ls[IQ2XS_NSUB]; uint16_t code[IQ2XS_NGROUP];
|
| 5915 |
+
for (int64_t blk = 0; blk < n_blocks; blk++) {
|
| 5916 |
+
if (bpr > 0 && (blk % bpr) == 0) rolling = 0.0f;
|
| 5917 |
+
const float *x = weights + blk * QK_K;
|
| 5918 |
+
BlockIQ2XS *b = &output[blk];
|
| 5919 |
+
float d = gguf_fp16_to_fp32(b->d);
|
| 5920 |
+
if (d > 0.0f) {
|
| 5921 |
+
float sigma2 = 0.0f;
|
| 5922 |
+
for (int i = 0; i < QK_K; i++) sigma2 += x[i] * x[i];
|
| 5923 |
+
sigma2 /= (float)QK_K;
|
| 5924 |
+
for (int i = 0; i < QK_K; i++) {
|
| 5925 |
+
float base = imat_importance ? imat_importance[blk * QK_K + i] : 1.0f;
|
| 5926 |
+
w[i] = (wmode == 1) ? base * sqrtf(sigma2 + x[i] * x[i]) : base;
|
| 5927 |
+
}
|
| 5928 |
+
for (int ib = 0; ib < IQ2XS_NSUB; ib++)
|
| 5929 |
+
ls[ib] = (ib & 1) ? (b->scales[ib >> 1] >> 4) : (b->scales[ib >> 1] & 0xF);
|
| 5930 |
+
memcpy(code, b->qs, sizeof(code));
|
| 5931 |
+
iq2xs_shape_block(x, w, d, ls, code, g_hex_dc_decay * rolling);
|
| 5932 |
+
memcpy(b->qs, code, sizeof(code));
|
| 5933 |
+
}
|
| 5934 |
+
iq2xs_dequant_block(b, deq);
|
| 5935 |
+
float r = 0.0f;
|
| 5936 |
+
for (int i = 0; i < QK_K; i++) r += x[i] - deq[i];
|
| 5937 |
+
rolling = r;
|
| 5938 |
+
}
|
| 5939 |
+
}
|
| 5940 |
+
|
| 5941 |
+
/* 6. exact reconstruction SSE */
|
| 5942 |
+
double tot = 0.0;
|
| 5943 |
+
#pragma omp parallel for reduction(+:tot)
|
| 5944 |
+
for (int64_t blk = 0; blk < n_blocks; blk++) {
|
| 5945 |
+
float deq[QK_K];
|
| 5946 |
+
iq2xs_dequant_block(&output[blk], deq);
|
| 5947 |
+
const float *x = weights + blk * QK_K;
|
| 5948 |
+
for (int i = 0; i < QK_K; i++) { double e = x[i] - deq[i]; tot += e * e; }
|
| 5949 |
+
}
|
| 5950 |
+
if (out_total_error) *out_total_error = (float)tot;
|
| 5951 |
+
if (verbose)
|
| 5952 |
+
printf(" [IQ2_XSΒ·Sieve] blocks=%lld rmse=%.4e\n", (long long)n_blocks,
|
| 5953 |
+
sqrt(tot / (double)n_elements));
|
| 5954 |
+
}
|
| 5955 |
+
|
| 5956 |
/* βββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββ
|
| 5957 |
* LIBRARY API β Exported functions for Python ctypes integration
|
| 5958 |
*
|
|
|
|
| 6059 |
imat_importance, verbose);
|
| 6060 |
}
|
| 6061 |
|
| 6062 |
+
/* IQ2_XS (74 bytes / 256 weights). Row-aware residual carry like Q2_K. */
|
| 6063 |
+
int hexstate_iq2xs_block_bytes(void) { return (int)sizeof(BlockIQ2XS); }
|
| 6064 |
+
int hexstate_iq2xs_block_elements(void) { return QK_K; }
|
| 6065 |
+
|
| 6066 |
+
void hexstate_quantize_tensor_iq2_xs_hpc(const float *weights, int64_t n_elements,
|
| 6067 |
+
void *output, float *out_error,
|
| 6068 |
+
const float *imat_importance, int verbose,
|
| 6069 |
+
int64_t row_width)
|
| 6070 |
+
{
|
| 6071 |
+
hexstate_init();
|
| 6072 |
+
quantize_tensor_iq2_xs_hpc(weights, n_elements, (BlockIQ2XS *)output,
|
| 6073 |
+
out_error, imat_importance, verbose, row_width);
|
| 6074 |
+
}
|
| 6075 |
+
|
| 6076 |
+
void hexstate_dequant_iq2_xs(const void *blocks, int64_t n_blocks, float *out)
|
| 6077 |
+
{
|
| 6078 |
+
const BlockIQ2XS *b = (const BlockIQ2XS *)blocks;
|
| 6079 |
+
for (int64_t i = 0; i < n_blocks; i++)
|
| 6080 |
+
iq2xs_dequant_block(&b[i], out + i * QK_K);
|
| 6081 |
+
}
|
| 6082 |
+
|
| 6083 |
#ifndef HEXSTATE_LIBRARY
|
| 6084 |
/* βββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββ
|
| 6085 |
* MAIN
|
hexstate_requantize.py
CHANGED
|
@@ -114,8 +114,44 @@ def _load_hexstate_lib():
|
|
| 114 |
lib.hexstate_set_spectral_params.restype = None
|
| 115 |
lib.hexstate_set_spectral_params.argtypes = [
|
| 116 |
ctypes.c_float, ctypes.c_float, ctypes.c_float]
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 117 |
|
| 118 |
lib.hexstate_init()
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 119 |
_HEXSTATE_LIB = lib
|
| 120 |
return lib
|
| 121 |
except Exception as e:
|
|
@@ -198,12 +234,20 @@ def read_imatrix(path):
|
|
| 198 |
base = name[:-len('.counts')]
|
| 199 |
counts_data[base] = data
|
| 200 |
|
| 201 |
-
#
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 202 |
for base_name in sum2_data:
|
| 203 |
in_sum2 = sum2_data[base_name]
|
| 204 |
count = counts_data.get(base_name, np.array([1.0]))[0]
|
| 205 |
if count > 0:
|
| 206 |
-
importance =
|
|
|
|
|
|
|
| 207 |
else:
|
| 208 |
importance = np.ones_like(in_sum2)
|
| 209 |
mean = importance.mean()
|
|
@@ -304,24 +348,31 @@ GGML_TYPE_F16 = 1
|
|
| 304 |
GGML_TYPE_Q4_0 = 2
|
| 305 |
GGML_TYPE_Q8_0 = 8
|
| 306 |
GGML_TYPE_Q2_K = 10
|
|
|
|
| 307 |
GGML_TYPE_BF16 = 30
|
| 308 |
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 309 |
TYPE_NAME = {
|
| 310 |
0: "F32", 1: "F16", 2: "Q4_0", 3: "Q4_1", 6: "Q5_0", 7: "Q5_1",
|
| 311 |
8: "Q8_0", 9: "Q8_1", 10: "Q2_K", 11: "Q3_K", 12: "Q4_K",
|
| 312 |
-
13: "Q5_K", 14: "Q6_K", 15: "Q8_K", 30: "BF16",
|
| 313 |
}
|
| 314 |
|
| 315 |
# Block sizes and byte sizes for each type
|
| 316 |
TYPE_BLOCK_SIZE = {
|
| 317 |
0: 1, 1: 1, 2: 32, 3: 32, 6: 32, 7: 32,
|
| 318 |
8: 32, 9: 32, 10: 256, 11: 256, 12: 256,
|
| 319 |
-
13: 256, 14: 256, 15: 256, 30: 1,
|
| 320 |
}
|
| 321 |
TYPE_BLOCK_BYTES = {
|
| 322 |
0: 4, 1: 2, 2: 18, 3: 20, 6: 20, 7: 22,
|
| 323 |
8: 34, 9: 36, 10: 84, 11: 110, 12: 144,
|
| 324 |
-
13: 176, 14: 210, 15: 292, 30: 2,
|
| 325 |
}
|
| 326 |
|
| 327 |
|
|
@@ -784,8 +835,42 @@ def _copy_bytes(fin, fout, abs_offset, n_bytes):
|
|
| 784 |
return written
|
| 785 |
|
| 786 |
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 787 |
def _stream_quantize(fin, fout, ti, abs_offset, kind, imatrix_data, use_hpc):
|
| 788 |
-
"""Quantize one tensor in row chunks. kind: 'q2k' | 'q4' | 'q8'.
|
| 789 |
Returns (n_out_bytes, rmse_or_None, sigma_or_None).
|
| 790 |
"""
|
| 791 |
ttype = ti['type']
|
|
@@ -798,7 +883,7 @@ def _stream_quantize(fin, fout, ti, abs_offset, kind, imatrix_data, use_hpc):
|
|
| 798 |
|
| 799 |
d0 = int(ti['dims'][0])
|
| 800 |
n_rows = int(ti['n_elements']) // d0
|
| 801 |
-
align = QK_K if kind
|
| 802 |
if d0 % align != 0:
|
| 803 |
raise ValueError(f'{ti["name"]} dim0={d0} not aligned to {align}')
|
| 804 |
|
|
@@ -816,7 +901,15 @@ def _stream_quantize(fin, fout, ti, abs_offset, kind, imatrix_data, use_hpc):
|
|
| 816 |
total_ss += float(np.vdot(f32, f32))
|
| 817 |
total_n += n_valid
|
| 818 |
|
| 819 |
-
if kind == '
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 820 |
if use_hpc:
|
| 821 |
qbytes, n_blocks = quantize_tensor_q2k_hpc(
|
| 822 |
f32, opt_mode=2, importance=imp, row_width=d0)
|
|
@@ -981,16 +1074,29 @@ def should_quantize(name, n_dims, dims, tied_embeddings=False):
|
|
| 981 |
def main():
|
| 982 |
if len(sys.argv) < 3:
|
| 983 |
print("Usage: python3 hexstate_requantize.py <input.gguf> <output.gguf>"
|
| 984 |
-
" [--keep-metadata] [--imatrix FILE] [--keep-embd] [--q2all]")
|
|
|
|
| 985 |
print(" HEX_CHUNK_ELEMS max f32 elements per tensor chunk (default 2000000)")
|
|
|
|
|
|
|
|
|
|
|
|
|
| 986 |
sys.exit(1)
|
| 987 |
|
|
|
|
| 988 |
input_path = sys.argv[1]
|
| 989 |
output_path = sys.argv[2]
|
| 990 |
keep_metadata = '--keep-metadata' in sys.argv
|
| 991 |
quantize_none = '--quantize-none' in sys.argv
|
| 992 |
q2all = '--q2all' in sys.argv
|
| 993 |
keep_embd = '--keep-embd' in sys.argv # keep tied embedding at source precision instead of Q8_0
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 994 |
|
| 995 |
# Check for imatrix
|
| 996 |
imatrix_data = None
|
|
@@ -1006,11 +1112,19 @@ def main():
|
|
| 1006 |
|
| 1007 |
# Check for HPC C library
|
| 1008 |
use_hpc = _load_hexstate_lib() is not None
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1009 |
|
| 1010 |
print()
|
| 1011 |
print(" βββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββοΏ½οΏ½")
|
| 1012 |
print(" β HExState GGUF Re-Quantizer β")
|
| 1013 |
-
print(" β GGUF β
|
|
|
|
|
|
|
| 1014 |
if q2all:
|
| 1015 |
print(" β Mode: --q2all ALL eligible tensors β Q2_K (test mode) β")
|
| 1016 |
if use_hpc and imatrix_data:
|
|
@@ -1178,9 +1292,9 @@ def main():
|
|
| 1178 |
out_size = n_blocks * 18
|
| 1179 |
print(f" Q4_0: {ti['name']} (dims[0]={dim0})")
|
| 1180 |
elif quant_plan[i] is True and q2k_row_compatible(ti['n_dims'], ti['dims']):
|
| 1181 |
-
out_type =
|
| 1182 |
n_blocks = ti['n_elements'] // QK_K
|
| 1183 |
-
out_size = n_blocks *
|
| 1184 |
else:
|
| 1185 |
out_type = ti['type']
|
| 1186 |
out_size = ti['data_size']
|
|
@@ -1210,8 +1324,8 @@ def main():
|
|
| 1210 |
else:
|
| 1211 |
for key, vtype, raw_value in kv_pairs:
|
| 1212 |
if key == 'general.file_type' and vtype == 4: # UINT32
|
| 1213 |
-
#
|
| 1214 |
-
updated_kv.append((key, vtype, struct.pack('<I',
|
| 1215 |
elif key == 'general.quantization_version' and vtype == 4:
|
| 1216 |
updated_kv.append((key, vtype, struct.pack('<I', 2)))
|
| 1217 |
elif key == 'tokenizer.ggml.token_type' and vtype == 9:
|
|
@@ -1369,14 +1483,14 @@ def main():
|
|
| 1369 |
|
| 1370 |
elif plan:
|
| 1371 |
nbytes, rmse, sigma = _stream_quantize(
|
| 1372 |
-
fin, fout, ti, abs_offset,
|
| 1373 |
if rmse is not None:
|
| 1374 |
q2k_rmse_sum += rmse
|
| 1375 |
q2k_tensor_count += 1
|
| 1376 |
-
print(f"\n [
|
| 1377 |
f" Ο={sigma:.4f} rel={rmse / max(sigma, 1e-30):.4f}")
|
| 1378 |
else:
|
| 1379 |
-
print(f"\n [
|
| 1380 |
quant_count += 1
|
| 1381 |
total_quant_bytes += nbytes
|
| 1382 |
|
|
@@ -1402,16 +1516,16 @@ def main():
|
|
| 1402 |
print(" ββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββ")
|
| 1403 |
print(" β RE-QUANTIZATION SUMMARY β")
|
| 1404 |
print(" β βββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββ£")
|
| 1405 |
-
print(f" β Tensors quantized (
|
| 1406 |
print(f" β Tensors kept as-is: {total_keep:<33d} β")
|
| 1407 |
-
print(f" β
|
| 1408 |
print(f" β Kept data: {total_keep_bytes:>12,} bytes ({total_keep_bytes/1024**2:>7.1f} MB) β")
|
| 1409 |
print(f" β Original size: {file_size:>12,} bytes ({file_size/1024**3:>7.2f} GB) β")
|
| 1410 |
print(f" β Output size: {final_size:>12,} bytes ({final_size/1024**3:>7.2f} GB) β")
|
| 1411 |
print(f" β Compression: {compression:>42.1f}x β")
|
| 1412 |
if q2k_tensor_count > 0:
|
| 1413 |
mean_rmse = q2k_rmse_sum / q2k_tensor_count
|
| 1414 |
-
print(f" β Mean
|
| 1415 |
print(f" β Total time: {elapsed:>39.1f} sec β")
|
| 1416 |
print(" ββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββ")
|
| 1417 |
print()
|
|
|
|
| 114 |
lib.hexstate_set_spectral_params.restype = None
|
| 115 |
lib.hexstate_set_spectral_params.argtypes = [
|
| 116 |
ctypes.c_float, ctypes.c_float, ctypes.c_float]
|
| 117 |
+
if hasattr(lib, 'hexstate_set_sse_budget'):
|
| 118 |
+
lib.hexstate_set_sse_budget.restype = None
|
| 119 |
+
lib.hexstate_set_sse_budget.argtypes = [ctypes.c_float]
|
| 120 |
+
|
| 121 |
+
# IQ2_XS (E8 codebook) quantizer + dequant
|
| 122 |
+
if hasattr(lib, 'hexstate_quantize_tensor_iq2_xs_hpc'):
|
| 123 |
+
lib.hexstate_quantize_tensor_iq2_xs_hpc.restype = None
|
| 124 |
+
lib.hexstate_quantize_tensor_iq2_xs_hpc.argtypes = [
|
| 125 |
+
ctypes.POINTER(ctypes.c_float), # weights
|
| 126 |
+
ctypes.c_int64, # n_elements
|
| 127 |
+
ctypes.c_void_p, # output
|
| 128 |
+
ctypes.POINTER(ctypes.c_float), # out_error
|
| 129 |
+
ctypes.POINTER(ctypes.c_float), # imat_importance (can be NULL)
|
| 130 |
+
ctypes.c_int, # verbose
|
| 131 |
+
ctypes.c_int64, # row_width
|
| 132 |
+
]
|
| 133 |
+
lib.hexstate_dequant_iq2_xs.restype = None
|
| 134 |
+
lib.hexstate_dequant_iq2_xs.argtypes = [
|
| 135 |
+
ctypes.c_void_p, ctypes.c_int64, ctypes.POINTER(ctypes.c_float)]
|
| 136 |
|
| 137 |
lib.hexstate_init()
|
| 138 |
+
dc_l = os.environ.get('HEX_DC_LAMBDA')
|
| 139 |
+
vw_l = os.environ.get('HEX_VW_LAMBDA')
|
| 140 |
+
dc_d = os.environ.get('HEX_DC_DECAY')
|
| 141 |
+
if hasattr(lib, 'hexstate_set_spectral_params') and (dc_l or vw_l or dc_d):
|
| 142 |
+
lib.hexstate_set_spectral_params(
|
| 143 |
+
ctypes.c_float(float(dc_l) if dc_l else 1.0),
|
| 144 |
+
ctypes.c_float(float(vw_l) if vw_l else 1.0),
|
| 145 |
+
ctypes.c_float(float(dc_d) if dc_d else 0.85),
|
| 146 |
+
)
|
| 147 |
+
# Relative SSE the DC/vesica shaper may spend. Q2_K default 5e-4;
|
| 148 |
+
# IQ2_XS has no per-sub-block offset, so cancelling DC means swapping
|
| 149 |
+
# codewords β it needs ~2e-2 (β +0.2% RMSE) to be effective.
|
| 150 |
+
budget = os.environ.get('HEX_SSE_BUDGET')
|
| 151 |
+
if budget is None and LOWBIT_FORMAT == 'iq2xs':
|
| 152 |
+
budget = '0.02'
|
| 153 |
+
if budget is not None and hasattr(lib, 'hexstate_set_sse_budget'):
|
| 154 |
+
lib.hexstate_set_sse_budget(ctypes.c_float(float(budget)))
|
| 155 |
_HEXSTATE_LIB = lib
|
| 156 |
return lib
|
| 157 |
except Exception as e:
|
|
|
|
| 234 |
base = name[:-len('.counts')]
|
| 235 |
counts_data[base] = data
|
| 236 |
|
| 237 |
+
# Importance = in_sum2 / counts = E[aΒ²] per input column, i.e. the
|
| 238 |
+
# weight in Ξ£ E[a_iΒ²]Β·e_iΒ² (output-error variance). This is exactly
|
| 239 |
+
# ggml's quant_weights and what the legacy .dat branch below returns.
|
| 240 |
+
# The earlier sqrt() here under-weighted important columns: in a
|
| 241 |
+
# controlled splice A/B (SmolLM2, ffn_down only) it cost +4.7% PPL
|
| 242 |
+
# for Q2_K and +7.8% for IQ2_XS. HEX_IMAT_SQRT=1 restores it.
|
| 243 |
+
use_sqrt = os.environ.get('HEX_IMAT_SQRT', '0') == '1'
|
| 244 |
for base_name in sum2_data:
|
| 245 |
in_sum2 = sum2_data[base_name]
|
| 246 |
count = counts_data.get(base_name, np.array([1.0]))[0]
|
| 247 |
if count > 0:
|
| 248 |
+
importance = in_sum2 / count
|
| 249 |
+
if use_sqrt:
|
| 250 |
+
importance = np.sqrt(importance)
|
| 251 |
else:
|
| 252 |
importance = np.ones_like(in_sum2)
|
| 253 |
mean = importance.mean()
|
|
|
|
| 348 |
GGML_TYPE_Q4_0 = 2
|
| 349 |
GGML_TYPE_Q8_0 = 8
|
| 350 |
GGML_TYPE_Q2_K = 10
|
| 351 |
+
GGML_TYPE_IQ2_XS = 17
|
| 352 |
GGML_TYPE_BF16 = 30
|
| 353 |
|
| 354 |
+
IQ2_XS_BLOCK_BYTES = 74 # d(fp16) + 32Γu16 codes + 8 scale bytes
|
| 355 |
+
|
| 356 |
+
# Low-bit target for the "Q2_K plan" tensors: 'q2k' (default) or 'iq2xs'
|
| 357 |
+
# (--iq2xs: E8 codebook, 2.3125 bpw, native llama.cpp decode).
|
| 358 |
+
LOWBIT_FORMAT = 'q2k'
|
| 359 |
+
|
| 360 |
TYPE_NAME = {
|
| 361 |
0: "F32", 1: "F16", 2: "Q4_0", 3: "Q4_1", 6: "Q5_0", 7: "Q5_1",
|
| 362 |
8: "Q8_0", 9: "Q8_1", 10: "Q2_K", 11: "Q3_K", 12: "Q4_K",
|
| 363 |
+
13: "Q5_K", 14: "Q6_K", 15: "Q8_K", 17: "IQ2_XS", 30: "BF16",
|
| 364 |
}
|
| 365 |
|
| 366 |
# Block sizes and byte sizes for each type
|
| 367 |
TYPE_BLOCK_SIZE = {
|
| 368 |
0: 1, 1: 1, 2: 32, 3: 32, 6: 32, 7: 32,
|
| 369 |
8: 32, 9: 32, 10: 256, 11: 256, 12: 256,
|
| 370 |
+
13: 256, 14: 256, 15: 256, 17: 256, 30: 1,
|
| 371 |
}
|
| 372 |
TYPE_BLOCK_BYTES = {
|
| 373 |
0: 4, 1: 2, 2: 18, 3: 20, 6: 20, 7: 22,
|
| 374 |
8: 34, 9: 36, 10: 84, 11: 110, 12: 144,
|
| 375 |
+
13: 176, 14: 210, 15: 292, 17: IQ2_XS_BLOCK_BYTES, 30: 2,
|
| 376 |
}
|
| 377 |
|
| 378 |
|
|
|
|
| 835 |
return written
|
| 836 |
|
| 837 |
|
| 838 |
+
def quantize_tensor_iq2xs_hpc(f32_data, importance=None, row_width=0):
|
| 839 |
+
"""IQ2_XS (E8 codebook, 74 B / 256 w) via the HPC C library.
|
| 840 |
+
Returns (bytes, n_blocks, dequant_f32)."""
|
| 841 |
+
lib = _load_hexstate_lib()
|
| 842 |
+
if lib is None or not hasattr(lib, 'hexstate_quantize_tensor_iq2_xs_hpc'):
|
| 843 |
+
raise RuntimeError('libhexstate_q2k.so lacks IQ2_XS support β rebuild')
|
| 844 |
+
f32 = np.ascontiguousarray(f32_data, dtype=np.float32).reshape(-1)
|
| 845 |
+
n = int(f32.size)
|
| 846 |
+
if n % QK_K != 0:
|
| 847 |
+
raise ValueError(f'IQ2_XS needs a multiple of {QK_K} elements, got {n}')
|
| 848 |
+
n_blocks = n // QK_K
|
| 849 |
+
out = np.zeros(n_blocks * IQ2_XS_BLOCK_BYTES, dtype=np.uint8)
|
| 850 |
+
err = ctypes.c_float(0.0)
|
| 851 |
+
imat_ptr = None
|
| 852 |
+
if importance is not None:
|
| 853 |
+
imat_c = np.ascontiguousarray(importance, dtype=np.float32).reshape(-1)
|
| 854 |
+
if imat_c.size == n:
|
| 855 |
+
imat_ptr = imat_c.ctypes.data_as(ctypes.POINTER(ctypes.c_float))
|
| 856 |
+
lib.hexstate_quantize_tensor_iq2_xs_hpc(
|
| 857 |
+
f32.ctypes.data_as(ctypes.POINTER(ctypes.c_float)),
|
| 858 |
+
ctypes.c_int64(n),
|
| 859 |
+
out.ctypes.data_as(ctypes.c_void_p),
|
| 860 |
+
ctypes.byref(err),
|
| 861 |
+
imat_ptr,
|
| 862 |
+
ctypes.c_int(0),
|
| 863 |
+
ctypes.c_int64(int(row_width)),
|
| 864 |
+
)
|
| 865 |
+
deq = np.zeros(n, dtype=np.float32)
|
| 866 |
+
lib.hexstate_dequant_iq2_xs(
|
| 867 |
+
out.ctypes.data_as(ctypes.c_void_p), ctypes.c_int64(n_blocks),
|
| 868 |
+
deq.ctypes.data_as(ctypes.POINTER(ctypes.c_float)))
|
| 869 |
+
return out.tobytes(), n_blocks, deq
|
| 870 |
+
|
| 871 |
+
|
| 872 |
def _stream_quantize(fin, fout, ti, abs_offset, kind, imatrix_data, use_hpc):
|
| 873 |
+
"""Quantize one tensor in row chunks. kind: 'q2k' | 'iq2xs' | 'q4' | 'q8'.
|
| 874 |
Returns (n_out_bytes, rmse_or_None, sigma_or_None).
|
| 875 |
"""
|
| 876 |
ttype = ti['type']
|
|
|
|
| 883 |
|
| 884 |
d0 = int(ti['dims'][0])
|
| 885 |
n_rows = int(ti['n_elements']) // d0
|
| 886 |
+
align = QK_K if kind in ('q2k', 'iq2xs') else 32
|
| 887 |
if d0 % align != 0:
|
| 888 |
raise ValueError(f'{ti["name"]} dim0={d0} not aligned to {align}')
|
| 889 |
|
|
|
|
| 901 |
total_ss += float(np.vdot(f32, f32))
|
| 902 |
total_n += n_valid
|
| 903 |
|
| 904 |
+
if kind == 'iq2xs':
|
| 905 |
+
qbytes, n_blocks, deq = quantize_tensor_iq2xs_hpc(
|
| 906 |
+
f32, importance=imp, row_width=d0)
|
| 907 |
+
fout.write(qbytes)
|
| 908 |
+
written += len(qbytes)
|
| 909 |
+
diff = f32.reshape(-1)[:n_valid] - deq[:n_valid]
|
| 910 |
+
total_se += float(np.sum(diff ** 2))
|
| 911 |
+
del qbytes, deq
|
| 912 |
+
elif kind == 'q2k':
|
| 913 |
if use_hpc:
|
| 914 |
qbytes, n_blocks = quantize_tensor_q2k_hpc(
|
| 915 |
f32, opt_mode=2, importance=imp, row_width=d0)
|
|
|
|
| 1074 |
def main():
|
| 1075 |
if len(sys.argv) < 3:
|
| 1076 |
print("Usage: python3 hexstate_requantize.py <input.gguf> <output.gguf>"
|
| 1077 |
+
" [--keep-metadata] [--imatrix FILE] [--keep-embd] [--q2all] [--iq2xs]")
|
| 1078 |
+
print(" --iq2xs low-bit tensors β IQ2_XS (E8 codebook, 2.3125 bpw) instead of Q2_K")
|
| 1079 |
print(" HEX_CHUNK_ELEMS max f32 elements per tensor chunk (default 2000000)")
|
| 1080 |
+
print(" HEX_DC_LAMBDA DC residual weight (default 1)")
|
| 1081 |
+
print(" HEX_VW_LAMBDA vesica weight (default 1)")
|
| 1082 |
+
print(" HEX_DC_DECAY rolling residual carry 0..1 (default 0.85)")
|
| 1083 |
+
print(" HEX_SSE_BUDGET relative SSE the shaper may spend (Q2_K 5e-4, IQ2_XS 2e-2)")
|
| 1084 |
sys.exit(1)
|
| 1085 |
|
| 1086 |
+
global LOWBIT_FORMAT
|
| 1087 |
input_path = sys.argv[1]
|
| 1088 |
output_path = sys.argv[2]
|
| 1089 |
keep_metadata = '--keep-metadata' in sys.argv
|
| 1090 |
quantize_none = '--quantize-none' in sys.argv
|
| 1091 |
q2all = '--q2all' in sys.argv
|
| 1092 |
keep_embd = '--keep-embd' in sys.argv # keep tied embedding at source precision instead of Q8_0
|
| 1093 |
+
if '--iq2xs' in sys.argv:
|
| 1094 |
+
LOWBIT_FORMAT = 'iq2xs'
|
| 1095 |
+
lowbit_kind = LOWBIT_FORMAT # 'q2k' | 'iq2xs'
|
| 1096 |
+
lowbit_name = 'IQ2_XS' if lowbit_kind == 'iq2xs' else 'Q2_K'
|
| 1097 |
+
lowbit_type = GGML_TYPE_IQ2_XS if lowbit_kind == 'iq2xs' else GGML_TYPE_Q2_K
|
| 1098 |
+
lowbit_bytes = IQ2_XS_BLOCK_BYTES if lowbit_kind == 'iq2xs' else 84
|
| 1099 |
+
lowbit_file_type = 20 if lowbit_kind == 'iq2xs' else 10 # LLAMA_FTYPE_MOSTLY_*
|
| 1100 |
|
| 1101 |
# Check for imatrix
|
| 1102 |
imatrix_data = None
|
|
|
|
| 1112 |
|
| 1113 |
# Check for HPC C library
|
| 1114 |
use_hpc = _load_hexstate_lib() is not None
|
| 1115 |
+
if lowbit_kind == 'iq2xs':
|
| 1116 |
+
lib = _load_hexstate_lib()
|
| 1117 |
+
if lib is None or not hasattr(lib, 'hexstate_quantize_tensor_iq2_xs_hpc'):
|
| 1118 |
+
print(" ERROR: --iq2xs needs libhexstate_q2k.so with IQ2_XS support "
|
| 1119 |
+
"(make -f makefile.quantize.c)")
|
| 1120 |
+
sys.exit(1)
|
| 1121 |
|
| 1122 |
print()
|
| 1123 |
print(" βββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββοΏ½οΏ½")
|
| 1124 |
print(" β HExState GGUF Re-Quantizer β")
|
| 1125 |
+
print(f" β GGUF β {lowbit_name:6s} GGUF with metadata passthrough β")
|
| 1126 |
+
if lowbit_kind == 'iq2xs':
|
| 1127 |
+
print(" β Low-bit: IQ2_XS E8 codebook Β· fold/DC shaping Β· 2.3125 bpw β")
|
| 1128 |
if q2all:
|
| 1129 |
print(" β Mode: --q2all ALL eligible tensors β Q2_K (test mode) β")
|
| 1130 |
if use_hpc and imatrix_data:
|
|
|
|
| 1292 |
out_size = n_blocks * 18
|
| 1293 |
print(f" Q4_0: {ti['name']} (dims[0]={dim0})")
|
| 1294 |
elif quant_plan[i] is True and q2k_row_compatible(ti['n_dims'], ti['dims']):
|
| 1295 |
+
out_type = lowbit_type
|
| 1296 |
n_blocks = ti['n_elements'] // QK_K
|
| 1297 |
+
out_size = n_blocks * lowbit_bytes
|
| 1298 |
else:
|
| 1299 |
out_type = ti['type']
|
| 1300 |
out_size = ti['data_size']
|
|
|
|
| 1324 |
else:
|
| 1325 |
for key, vtype, raw_value in kv_pairs:
|
| 1326 |
if key == 'general.file_type' and vtype == 4: # UINT32
|
| 1327 |
+
# LLAMA_FTYPE_MOSTLY_Q2_K = 10, MOSTLY_IQ2_XS = 20
|
| 1328 |
+
updated_kv.append((key, vtype, struct.pack('<I', lowbit_file_type)))
|
| 1329 |
elif key == 'general.quantization_version' and vtype == 4:
|
| 1330 |
updated_kv.append((key, vtype, struct.pack('<I', 2)))
|
| 1331 |
elif key == 'tokenizer.ggml.token_type' and vtype == 9:
|
|
|
|
| 1483 |
|
| 1484 |
elif plan:
|
| 1485 |
nbytes, rmse, sigma = _stream_quantize(
|
| 1486 |
+
fin, fout, ti, abs_offset, lowbit_kind, imatrix_data, use_hpc)
|
| 1487 |
if rmse is not None:
|
| 1488 |
q2k_rmse_sum += rmse
|
| 1489 |
q2k_tensor_count += 1
|
| 1490 |
+
print(f"\n [{lowbit_name}] {ti['name'][:50]} RMSE={rmse:.6e}"
|
| 1491 |
f" Ο={sigma:.4f} rel={rmse / max(sigma, 1e-30):.4f}")
|
| 1492 |
else:
|
| 1493 |
+
print(f"\n [{lowbit_name}] {ti['name'][:55]} RMSE=n/a")
|
| 1494 |
quant_count += 1
|
| 1495 |
total_quant_bytes += nbytes
|
| 1496 |
|
|
|
|
| 1516 |
print(" ββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββ")
|
| 1517 |
print(" β RE-QUANTIZATION SUMMARY β")
|
| 1518 |
print(" β βββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββ£")
|
| 1519 |
+
print(f" β Tensors quantized ({lowbit_name:6s}): {quant_count:<31d} β")
|
| 1520 |
print(f" β Tensors kept as-is: {total_keep:<33d} β")
|
| 1521 |
+
print(f" β {lowbit_name:6s} data: {total_quant_bytes:>12,} bytes ({total_quant_bytes/1024**2:>7.1f} MB) β")
|
| 1522 |
print(f" β Kept data: {total_keep_bytes:>12,} bytes ({total_keep_bytes/1024**2:>7.1f} MB) β")
|
| 1523 |
print(f" β Original size: {file_size:>12,} bytes ({file_size/1024**3:>7.2f} GB) β")
|
| 1524 |
print(f" β Output size: {final_size:>12,} bytes ({final_size/1024**3:>7.2f} GB) β")
|
| 1525 |
print(f" β Compression: {compression:>42.1f}x β")
|
| 1526 |
if q2k_tensor_count > 0:
|
| 1527 |
mean_rmse = q2k_rmse_sum / q2k_tensor_count
|
| 1528 |
+
print(f" β Mean {lowbit_name:6s} RMSE: {mean_rmse:>12.6e} β")
|
| 1529 |
print(f" β Total time: {elapsed:>39.1f} sec β")
|
| 1530 |
print(" ββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββ")
|
| 1531 |
print()
|