Not a member of Pastebin yet?
Sign Up,
it unlocks many cool features!
- // Helper: set scale and min in packed format
- static inline void set_scale_min_k4(int j, uint8_t * GGML_RESTRICT q, uint8_t d, uint8_t m) {
- assert(d < 64 && m < 64);
- if (j < 4) {
- q[j] = (q[j] & 0xC0) | (d & 0x3F);
- q[j + 4] = (q[j + 4] & 0xC0) | (m & 0x3F);
- } else {
- const int j2 = j - 4;
- q[j2] = (q[j2] & 0x3F) | ((d & 0x30) << 2);
- q[j + 4] = (d & 0x0F) | ((m & 0x0F) << 4);
- q[j] = (q[j] & 0x3F) | ((m & 0x30) << 2);
- }
- }
- #define DEBUG_LLOYD_MAX
- void quantize_row_q4_K_ref(const float * GGML_RESTRICT x, block_q4_K * GGML_RESTRICT y, int64_t k) {
- assert(k % QK_K == 0);
- const int max_iter = 10;
- const float epsilon = 1e-6f;
- const int nb = k / QK_K;
- const int num_subblocks = QK_K / 32;
- for (int i = 0; i < nb; i++) {
- memset(y[i].scales, 0, K_SCALE_SIZE);
- float scales[num_subblocks];
- float mins[num_subblocks];
- // Initialization: exhaustive search over all divisors and offsets
- // div=15..1, offset_bin=0..(15-div), total 120 combinations
- // offset_bin = which QAT bin the observed xmin maps to
- for (int j = 0; j < num_subblocks; j++) {
- float xmin = x[i*QK_K + j*32];
- float xmax = x[i*QK_K + j*32];
- for (int l = 1; l < 32; l++) {
- const float v = x[i*QK_K + j*32 + l];
- xmin = v < xmin ? v : xmin;
- xmax = v > xmax ? v : xmax;
- }
- float xrange = xmax - xmin;
- if (xrange == 0.0f) {
- scales[j] = 1.0f;
- mins[j] = xmin;
- continue;
- }
- int best_divisor = 15;
- int best_offset_bin = 0; // which bin xmin maps to
- float best_error = FLT_MAX;
- // Test all divisors from 15 down to 1
- for (int div = 15; div >= 1; div--) {
- // offset_bin can be 0 to (15 - div)
- // This ensures max_bin = offset_bin + div <= 15
- int max_offset_bin = 15 - div;
- for (int offset_bin = 0; offset_bin <= max_offset_bin; offset_bin++) {
- float scale = xrange / (float)div;
- // offset = xmin - offset_bin * scale
- // So when we compute qi = round((v - offset) / scale)
- // and v = xmin, we get qi = offset_bin
- float offset = xmin - offset_bin * scale;
- // Compute SSE for this (div, offset_bin) combination
- float error = 0.0f;
- for (int l = 0; l < 32; l++) {
- const float v = x[i*QK_K + j*32 + l];
- int q = nearest_int((v - offset) / scale);
- q = q < 0 ? 0 : (q > 15 ? 15 : q);
- float v_recon = offset + q * scale;
- float diff = v - v_recon;
- error += diff * diff;
- }
- if (error < best_error) {
- best_error = error;
- best_divisor = div;
- best_offset_bin = offset_bin;
- }
- }
- }
- scales[j] = xrange / (float)best_divisor;
- mins[j] = xmin - best_offset_bin * scales[j];
- #ifdef DEBUG_LLOYD_MAX
- printf("- [%i][%i] init: best_divisor=%d, best_offset_bin=%d, best_error=%.6f, scale=%.6f, min=%.6f\n",
- i % QK_K, j, best_divisor, best_offset_bin, (double)best_error, (double)scales[j], (double)mins[j]);
- #endif
- }
- // Initialize super-block scales
- float d = 0.0f;
- float dmin_abs = 0.0f;
- for (int j = 0; j < num_subblocks; j++) {
- d = scales[j] > d ? scales[j] : d;
- const float mins_abs = fabsf(mins[j]);
- dmin_abs = mins_abs > dmin_abs ? mins_abs : dmin_abs;
- }
- d = d / 63.0f;
- float dmin = dmin_abs / 63.0f;
- if (d == 0.0f) d = 1.0f;
- if (dmin == 0.0f) dmin = 1.0f;
- // Quantize initial sub-block scales and mins
- uint8_t sc[num_subblocks];
- uint8_t m[num_subblocks];
- for (int j = 0; j < num_subblocks; j++) {
- sc[j] = (uint8_t)(nearest_int(scales[j] / d));
- sc[j] = sc[j] > 63 ? 63 : sc[j];
- const int m_int = nearest_int(mins[j] / dmin);
- m[j] = (uint8_t)(m_int < 0 ? -m_int : m_int);
- m[j] = m[j] > 63 ? 63 : m[j];
- set_scale_min_k4(j, y[i].scales, sc[j], m[j]);
- }
- // Adjust dmin sign based on typical min values
- float avg_min = 0.0f;
- for (int j = 0; j < num_subblocks; j++) avg_min += mins[j];
- avg_min /= num_subblocks;
- if (avg_min > 0.0f) dmin = -dmin;
- // Temporary storage for 4-bit codes
- uint8_t q[QK_K];
- // Lloyd-Max iteration
- for (int iter = 0; iter < max_iter; iter++) {
- const float d_old = d;
- const float dmin_old = dmin;
- // Step 1: Assignment - quantize to 4-bit codes
- for (int j = 0; j < num_subblocks; j++) {
- const float scale = d * sc[j];
- const float offset = -dmin * m[j];
- if (scale == 0.0f) {
- for (int l = 0; l < 32; ++l) {
- q[j*32 + l] = 0;
- }
- continue;
- }
- for (int l = 0; l < 32; l++) {
- const float v = x[i*QK_K + j*32 + l];
- const int q_int = nearest_int((v - offset) / scale);
- q[j*32 + l] = (uint8_t)(q_int < 0 ? 0 : (q_int > 15 ? 15 : q_int));
- }
- }
- // Step 2: Update sub-block scales and mins (2D least squares per sub-block)
- for (int j = 0; j < num_subblocks; j++) {
- float sum_x = 0.0f;
- float sum_q = 0.0f;
- float sum_xq = 0.0f;
- float sum_qq = 0.0f;
- for (int l = 0; l < 32; l++) {
- const float xv = x[i*QK_K + j*32 + l];
- const float qv = (float)q[j*32 + l];
- sum_x += xv;
- sum_q += qv;
- sum_xq += xv * qv;
- sum_qq += qv * qv;
- }
- const float n = 32.0f;
- const float det = n * sum_qq - sum_q * sum_q;
- if (det > 0.0f) {
- const float a = (n * sum_xq - sum_x * sum_q) / det;
- const float b = (sum_x - a * sum_q) / n;
- if (a > 0.0f && d > 0.0f) {
- const int sc_new = nearest_int(a / d);
- sc[j] = (uint8_t)(sc_new < 0 ? 0 : (sc_new > 63 ? 63 : sc_new));
- }
- if (dmin != 0.0f) {
- const int m_new = nearest_int(-b / dmin);
- m[j] = (uint8_t)(m_new < 0 ? 0 : (m_new > 63 ? 63 : m_new));
- }
- set_scale_min_k4(j, y[i].scales, sc[j], m[j]);
- }
- }
- // Step 3: Update super-block scales (2D least squares across all sub-blocks)
- float A = 0.0f; // Σ(sc*q)²
- float B = 0.0f; // Σ(m*sc*q)
- float C = 0.0f; // Σm²
- float X_d = 0.0f; // Σ(x*sc*q)
- float X_m = 0.0f; // Σ(x*m)
- for (int j = 0; j < num_subblocks; j++) {
- float sum_sq = 0.0f;
- float sum_q = 0.0f;
- float sum_xq = 0.0f;
- float sum_x = 0.0f;
- for (int l = 0; l < 32; l++) {
- const float xv = x[i*QK_K + j*32 + l];
- const float qv = (float)q[j*32 + l];
- sum_sq += qv * qv;
- sum_q += qv;
- sum_xq += xv * qv;
- sum_x += xv;
- }
- const float sc_f = (float)sc[j];
- const float m_f = (float)m[j];
- A += sc_f * sc_f * sum_sq;
- B += m_f * sc_f * sum_q;
- C += m_f * m_f * 32.0f;
- X_d += sc_f * sum_xq;
- X_m += m_f * sum_x;
- }
- const float det = A * C - B * B;
- if (det > 0.0f) {
- const float d_new = (C * X_d - B * X_m) / det;
- const float dmin_new = (B * X_d - A * X_m) / det;
- if (d_new > 0.0f) {
- d = d_new;
- }
- if (dmin_new != 0.0f) {
- dmin = dmin_new;
- }
- }
- // Check convergence
- const float delta_d = fabsf(d - d_old);
- const float delta_dmin = fabsf(dmin - dmin_old);
- #ifdef DEBUG_LLOYD_MAX
- printf("- [%i] iter %i: delta_d=%.7f, delta_dmin=%.7f\n", i % QK_K, iter + 1, (double)delta_d, (double)delta_dmin);
- #endif
- if (delta_d < epsilon && delta_dmin < epsilon) {
- break;
- }
- }
- // Final assignment with converged parameters
- for (int j = 0; j < num_subblocks; j++) {
- const float scale = d * sc[j];
- const float offset = -dmin * m[j];
- for (int l = 0; l < 32; l++) {
- const float v = x[i*QK_K + j*32 + l];
- const int q_int = scale != 0.0f ? nearest_int((v - offset) / scale) : 0;
- q[j*32 + l] = (uint8_t)(q_int < 0 ? 0 : (q_int > 15 ? 15 : q_int));
- }
- }
- // Store final super-block scales
- y[i].d = GGML_FP32_TO_FP16(d);
- y[i].dmin = GGML_FP32_TO_FP16(dmin);
- // Pack 4-bit quantized values (layout expected by dequant)
- uint8_t *qs = y[i].qs;
- for (int base = 0, out = 0; base < QK_K; base += 64, out += 32) {
- for (int l = 0; l < 32; ++l) {
- qs[out + l] = (q[base + l] & 0x0F) | ((q[base + 32 + l] & 0x0F) << 4);
- }
- }
- // Dequantize and check error
- #ifdef DEBUG_LLOYD_MAX
- float y_dequant[QK_K];
- dequantize_row_q4_K(&y[i], y_dequant, QK_K);
- float sum_error = 0.0f;
- float sum_abs_error = 0.0f;
- for (int j = 0; j < QK_K; j++) {
- const float error = y_dequant[j] - x[i*QK_K + j];
- printf("- [%i][%i] final: original=%.6f, dequant=%.6f, error=%.6f\n", i % QK_K, j, (double)x[i*QK_K + j], (double)y_dequant[j], (double)error);
- sum_error += error;
- sum_abs_error += fabsf(error);
- }
- const float mean_error = sum_error / QK_K;
- const float mean_abs_error = sum_abs_error / QK_K;
- printf("- Mean error : %.6f\n", (double)mean_error);
- printf("- Mean absolute error: %.6f\n\n", (double)mean_abs_error);
- #endif
- }
- }
Add Comment
Please, Sign In to add comment