Guest User

Untitled

a guest
Sep 5th, 2026
20
0
Never
Not a member of Pastebin yet? Sign Up, it unlocks many cool features!
C 11.04 KB | None | 0 0
  1. // Helper: set scale and min in packed format
  2. static inline void set_scale_min_k4(int j, uint8_t * GGML_RESTRICT q, uint8_t d, uint8_t m) {
  3.     assert(d < 64 && m < 64);
  4.     if (j < 4) {
  5.         q[j] = (q[j] & 0xC0) | (d & 0x3F);
  6.         q[j + 4] = (q[j + 4] & 0xC0) | (m & 0x3F);
  7.     } else {
  8.         const int j2 = j - 4;
  9.         q[j2] = (q[j2] & 0x3F) | ((d & 0x30) << 2);
  10.         q[j + 4] = (d & 0x0F) | ((m & 0x0F) << 4);
  11.         q[j] = (q[j] & 0x3F) | ((m & 0x30) << 2);
  12.     }
  13. }
  14.  
  15. #define DEBUG_LLOYD_MAX
  16.  
  17. void quantize_row_q4_K_ref(const float * GGML_RESTRICT x, block_q4_K * GGML_RESTRICT y, int64_t k) {
  18.     assert(k % QK_K == 0);
  19.  
  20.     const int max_iter = 10;
  21.     const float epsilon = 1e-6f;
  22.     const int nb = k / QK_K;
  23.     const int num_subblocks = QK_K / 32;
  24.     for (int i = 0; i < nb; i++) {
  25.         memset(y[i].scales, 0, K_SCALE_SIZE);
  26.  
  27.         float scales[num_subblocks];
  28.         float mins[num_subblocks];
  29.  
  30.         // Initialization: exhaustive search over all divisors and offsets
  31.         // div=15..1, offset_bin=0..(15-div), total 120 combinations
  32.         // offset_bin = which QAT bin the observed xmin maps to
  33.         for (int j = 0; j < num_subblocks; j++) {
  34.             float xmin = x[i*QK_K + j*32];
  35.             float xmax = x[i*QK_K + j*32];
  36.  
  37.             for (int l = 1; l < 32; l++) {
  38.                 const float v = x[i*QK_K + j*32 + l];
  39.                 xmin = v < xmin ? v : xmin;
  40.                 xmax = v > xmax ? v : xmax;
  41.             }
  42.  
  43.             float xrange = xmax - xmin;
  44.             if (xrange == 0.0f) {
  45.                 scales[j] = 1.0f;
  46.                 mins[j] = xmin;
  47.                 continue;
  48.             }
  49.  
  50.             int best_divisor = 15;
  51.             int best_offset_bin = 0;  // which bin xmin maps to
  52.             float best_error = FLT_MAX;
  53.  
  54.             // Test all divisors from 15 down to 1
  55.             for (int div = 15; div >= 1; div--) {
  56.                 // offset_bin can be 0 to (15 - div)
  57.                 // This ensures max_bin = offset_bin + div <= 15
  58.                 int max_offset_bin = 15 - div;
  59.  
  60.                 for (int offset_bin = 0; offset_bin <= max_offset_bin; offset_bin++) {
  61.                     float scale = xrange / (float)div;
  62.                     // offset = xmin - offset_bin * scale
  63.                     // So when we compute qi = round((v - offset) / scale)
  64.                     // and v = xmin, we get qi = offset_bin
  65.                     float offset = xmin - offset_bin * scale;
  66.  
  67.                     // Compute SSE for this (div, offset_bin) combination
  68.                     float error = 0.0f;
  69.                     for (int l = 0; l < 32; l++) {
  70.                         const float v = x[i*QK_K + j*32 + l];
  71.                         int q = nearest_int((v - offset) / scale);
  72.                         q = q < 0 ? 0 : (q > 15 ? 15 : q);
  73.                         float v_recon = offset + q * scale;
  74.                         float diff = v - v_recon;
  75.                         error += diff * diff;
  76.                     }
  77.  
  78.                     if (error < best_error) {
  79.                         best_error = error;
  80.                         best_divisor = div;
  81.                         best_offset_bin = offset_bin;
  82.                     }
  83.                 }
  84.             }
  85.            
  86.             scales[j] = xrange / (float)best_divisor;
  87.             mins[j] = xmin - best_offset_bin * scales[j];
  88.            
  89.             #ifdef DEBUG_LLOYD_MAX
  90.             printf("- [%i][%i] init: best_divisor=%d, best_offset_bin=%d, best_error=%.6f, scale=%.6f, min=%.6f\n",
  91.                    i % QK_K, j, best_divisor, best_offset_bin, (double)best_error, (double)scales[j], (double)mins[j]);
  92.             #endif
  93.         }
  94.  
  95.         // Initialize super-block scales
  96.         float d = 0.0f;
  97.         float dmin_abs = 0.0f;
  98.         for (int j = 0; j < num_subblocks; j++) {
  99.             d = scales[j] > d ? scales[j] : d;
  100.             const float mins_abs = fabsf(mins[j]);
  101.             dmin_abs = mins_abs > dmin_abs ? mins_abs : dmin_abs;
  102.         }
  103.         d = d / 63.0f;
  104.         float dmin = dmin_abs / 63.0f;
  105.         if (d == 0.0f) d = 1.0f;
  106.         if (dmin == 0.0f) dmin = 1.0f;
  107.  
  108.         // Quantize initial sub-block scales and mins
  109.         uint8_t sc[num_subblocks];
  110.         uint8_t m[num_subblocks];
  111.         for (int j = 0; j < num_subblocks; j++) {
  112.             sc[j] = (uint8_t)(nearest_int(scales[j] / d));
  113.             sc[j] = sc[j] > 63 ? 63 : sc[j];
  114.  
  115.             const int m_int = nearest_int(mins[j] / dmin);
  116.             m[j] = (uint8_t)(m_int < 0 ? -m_int : m_int);
  117.             m[j] = m[j] > 63 ? 63 : m[j];
  118.  
  119.             set_scale_min_k4(j, y[i].scales, sc[j], m[j]);
  120.         }
  121.  
  122.         // Adjust dmin sign based on typical min values
  123.         float avg_min = 0.0f;
  124.         for (int j = 0; j < num_subblocks; j++) avg_min += mins[j];
  125.         avg_min /= num_subblocks;
  126.         if (avg_min > 0.0f) dmin = -dmin;
  127.  
  128.         // Temporary storage for 4-bit codes
  129.         uint8_t q[QK_K];
  130.  
  131.         // Lloyd-Max iteration
  132.         for (int iter = 0; iter < max_iter; iter++) {
  133.             const float d_old = d;
  134.             const float dmin_old = dmin;
  135.  
  136.             // Step 1: Assignment - quantize to 4-bit codes
  137.             for (int j = 0; j < num_subblocks; j++) {
  138.                 const float scale = d * sc[j];
  139.                 const float offset = -dmin * m[j];
  140.  
  141.                 if (scale == 0.0f) {
  142.                     for (int l = 0; l < 32; ++l) {
  143.                         q[j*32 + l] = 0;
  144.                     }
  145.                     continue;
  146.                 }
  147.  
  148.                 for (int l = 0; l < 32; l++) {
  149.                     const float v = x[i*QK_K + j*32 + l];
  150.                     const int q_int = nearest_int((v - offset) / scale);
  151.                     q[j*32 + l] = (uint8_t)(q_int < 0 ? 0 : (q_int > 15 ? 15 : q_int));
  152.                 }
  153.             }
  154.  
  155.             // Step 2: Update sub-block scales and mins (2D least squares per sub-block)
  156.             for (int j = 0; j < num_subblocks; j++) {
  157.                 float sum_x = 0.0f;
  158.                 float sum_q = 0.0f;
  159.                 float sum_xq = 0.0f;
  160.                 float sum_qq = 0.0f;
  161.  
  162.                 for (int l = 0; l < 32; l++) {
  163.                     const float xv = x[i*QK_K + j*32 + l];
  164.                     const float qv = (float)q[j*32 + l];
  165.                     sum_x += xv;
  166.                     sum_q += qv;
  167.                     sum_xq += xv * qv;
  168.                     sum_qq += qv * qv;
  169.                 }
  170.  
  171.                 const float n = 32.0f;
  172.                 const float det = n * sum_qq - sum_q * sum_q;
  173.  
  174.                 if (det > 0.0f) {
  175.                     const float a = (n * sum_xq - sum_x * sum_q) / det;
  176.                     const float b = (sum_x - a * sum_q) / n;
  177.  
  178.                     if (a > 0.0f && d > 0.0f) {
  179.                         const int sc_new = nearest_int(a / d);
  180.                         sc[j] = (uint8_t)(sc_new < 0 ? 0 : (sc_new > 63 ? 63 : sc_new));
  181.                     }
  182.  
  183.                     if (dmin != 0.0f) {
  184.                         const int m_new = nearest_int(-b / dmin);
  185.                         m[j] = (uint8_t)(m_new < 0 ? 0 : (m_new > 63 ? 63 : m_new));
  186.                     }
  187.  
  188.                     set_scale_min_k4(j, y[i].scales, sc[j], m[j]);
  189.                 }
  190.             }
  191.  
  192.             // Step 3: Update super-block scales (2D least squares across all sub-blocks)
  193.             float A = 0.0f;   // Σ(sc*q)²
  194.             float B = 0.0f;   // Σ(m*sc*q)
  195.             float C = 0.0f;   // Σm²
  196.             float X_d = 0.0f; // Σ(x*sc*q)
  197.             float X_m = 0.0f; // Σ(x*m)
  198.  
  199.             for (int j = 0; j < num_subblocks; j++) {
  200.                 float sum_sq = 0.0f;
  201.                 float sum_q = 0.0f;
  202.                 float sum_xq = 0.0f;
  203.                 float sum_x = 0.0f;
  204.  
  205.                 for (int l = 0; l < 32; l++) {
  206.                     const float xv = x[i*QK_K + j*32 + l];
  207.                     const float qv = (float)q[j*32 + l];
  208.                     sum_sq += qv * qv;
  209.                     sum_q += qv;
  210.                     sum_xq += xv * qv;
  211.                     sum_x += xv;
  212.                 }
  213.  
  214.                 const float sc_f = (float)sc[j];
  215.                 const float m_f = (float)m[j];
  216.  
  217.                 A += sc_f * sc_f * sum_sq;
  218.                 B += m_f * sc_f * sum_q;
  219.                 C += m_f * m_f * 32.0f;
  220.                 X_d += sc_f * sum_xq;
  221.                 X_m += m_f * sum_x;
  222.             }
  223.  
  224.             const float det = A * C - B * B;
  225.  
  226.             if (det > 0.0f) {
  227.                 const float d_new = (C * X_d - B * X_m) / det;
  228.                 const float dmin_new = (B * X_d - A * X_m) / det;
  229.  
  230.                 if (d_new > 0.0f) {
  231.                     d = d_new;
  232.                 }
  233.                 if (dmin_new != 0.0f) {
  234.                     dmin = dmin_new;
  235.                 }
  236.             }
  237.  
  238.             // Check convergence
  239.             const float delta_d = fabsf(d - d_old);
  240.             const float delta_dmin = fabsf(dmin - dmin_old);
  241.  
  242.             #ifdef DEBUG_LLOYD_MAX
  243.             printf("- [%i] iter %i: delta_d=%.7f, delta_dmin=%.7f\n", i % QK_K, iter + 1, (double)delta_d, (double)delta_dmin);
  244.             #endif
  245.            
  246.             if (delta_d < epsilon && delta_dmin < epsilon) {
  247.                 break;
  248.             }
  249.         }
  250.  
  251.         // Final assignment with converged parameters
  252.         for (int j = 0; j < num_subblocks; j++) {
  253.             const float scale = d * sc[j];
  254.             const float offset = -dmin * m[j];
  255.  
  256.             for (int l = 0; l < 32; l++) {
  257.                 const float v = x[i*QK_K + j*32 + l];
  258.                 const int q_int = scale != 0.0f ? nearest_int((v - offset) / scale) : 0;
  259.                 q[j*32 + l] = (uint8_t)(q_int < 0 ? 0 : (q_int > 15 ? 15 : q_int));
  260.             }
  261.         }
  262.  
  263.         // Store final super-block scales
  264.         y[i].d = GGML_FP32_TO_FP16(d);
  265.         y[i].dmin = GGML_FP32_TO_FP16(dmin);
  266.  
  267.         // Pack 4-bit quantized values (layout expected by dequant)
  268.         uint8_t *qs = y[i].qs;
  269.         for (int base = 0, out = 0; base < QK_K; base += 64, out += 32) {
  270.             for (int l = 0; l < 32; ++l) {
  271.                 qs[out + l] = (q[base + l] & 0x0F) | ((q[base + 32 + l] & 0x0F) << 4);
  272.             }
  273.         }
  274.  
  275.         // Dequantize and check error
  276.         #ifdef DEBUG_LLOYD_MAX
  277.         float y_dequant[QK_K];
  278.         dequantize_row_q4_K(&y[i], y_dequant, QK_K);
  279.         float sum_error = 0.0f;
  280.         float sum_abs_error = 0.0f;
  281.         for (int j = 0; j < QK_K; j++) {
  282.             const float error = y_dequant[j] - x[i*QK_K + j];
  283.             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);
  284.             sum_error += error;
  285.             sum_abs_error += fabsf(error);
  286.         }
  287.         const float mean_error = sum_error / QK_K;
  288.         const float mean_abs_error = sum_abs_error / QK_K;
  289.         printf("- Mean error         : %.6f\n", (double)mean_error);
  290.         printf("- Mean absolute error: %.6f\n\n", (double)mean_abs_error);
  291.         #endif
  292.     }
  293. }
Add Comment
Please, Sign In to add comment