| """GGUF k-quant packed slice-dequantization — torch, GPU-friendly, output-row slice. |
| |
| Reads packed GGUF k-quant raw bytes (uint8 [out, bytes_per_row]) and dequantizes |
| ONLY output rows [start:end] on-the-fly — for chunked matmul via QuantizedModule. |
| Weights stay packed in VRAM; only the requested row slice is unpacked to fp16/fp32. |
| |
| This is the memory-efficient path: a Q3_K_S 12 GB model stays ~12 GB in VRAM, |
| with only a small per-chunk dequant overhead (chunk_size rows at a time). |
| |
| Block layouts (from ggml-common.h, same as gguf_kquant_unpack.py): |
| Q8_0 (34 B): d(f16) + qs[32] (int8). y = d * q8 |
| Q4_K (144 B): d(f16) + dmin(f16) + scales[12] + qs[128] (4-bit). |
| y = d * sc * q4 - dmin * m (8 sub-blocks × 32, asymmetric) |
| Q5_K (176 B): d(f16) + dmin(f16) + scales[12] + qh[32] + qs[128] (4-bit). |
| y = d * sc * (q4 + 16*bit5) - dmin * m |
| Q6_K (210 B): ql[128] + qh[64] + scales[16] (int8) + d(f16). |
| y = d * sc * (q6 - 32) (16 sub-blocks × 16, symmetric) |
| |
| All operations are on torch tensors (view uint8/int8, bit ops via int16/int32 |
| intermediate) so they run on GPU when the buffer is on cuda. |
| """ |
|
|
| from __future__ import annotations |
|
|
| import torch |
|
|
| from agiws_neural_quant.kquant._gguf_bits import f16_view_to_f32, unpack_scales_k4 |
| from agiws_neural_quant.kquant.gguf_packed_q2q3 import ( |
| dequant_q2_k_packed_rows, dequant_q3_k_packed_rows, |
| ) |
| from agiws_neural_quant.kquant.iq import ( |
| dequant_iq2_xxs_packed_rows, |
| dequant_iq2_xs_packed, |
| dequant_iq2_s_packed, |
| dequant_iq3_xxs_packed, |
| dequant_iq3_s_packed, |
| dequant_iq1_s_packed, |
| dequant_iq1_m_packed, |
| dequant_iq4_nl_packed, |
| dequant_iq4_xs_packed, |
| ) |
|
|
|
|
| |
| |
| GGML_TYPE_Q8_0 = 8 |
| GGML_TYPE_Q4_K = 12 |
| GGML_TYPE_Q5_K = 13 |
| GGML_TYPE_Q6_K = 14 |
| GGML_TYPE_Q2_K = 10 |
| GGML_TYPE_Q3_K = 11 |
| GGML_TYPE_IQ2_XXS = 16 |
| GGML_TYPE_IQ2_XS = 17 |
| GGML_TYPE_IQ3_XXS = 18 |
| GGML_TYPE_IQ1_S = 19 |
| GGML_TYPE_IQ4_NL = 20 |
| GGML_TYPE_IQ3_S = 21 |
| GGML_TYPE_IQ2_S = 22 |
| GGML_TYPE_IQ4_XS = 23 |
| GGML_TYPE_IQ1_M = 29 |
|
|
|
|
| |
| def _f16_view_to_f32(u8: torch.Tensor) -> torch.Tensor: |
| return f16_view_to_f32(u8) |
|
|
|
|
| def _unpack_scales_k4(scales: torch.Tensor): |
| return unpack_scales_k4(scales) |
|
|
|
|
| def dequant_q8_0_packed_rows( |
| raw: torch.Tensor, start: int, end: int, cols: int |
| ) -> torch.Tensor: |
| """Dequant Q8_0 rows [start:end]. raw: [out, bytes_per_row] uint8.""" |
| block = 32 |
| elem = 34 |
| chunk = raw[start:end].to(torch.int32) |
| cr = chunk.shape[0] |
| n_blocks = (cols + block - 1) // block |
| chunk = chunk.reshape(cr, n_blocks, elem) |
| d = _f16_view_to_f32(chunk[..., 0:2]) |
| qs = chunk[..., 2:34].to(torch.int8).to(torch.float32) |
| y = qs * d.unsqueeze(2) |
| y = y.reshape(cr, n_blocks * block)[:, :cols] |
| return y |
|
|
|
|
| def dequant_q4_k_packed_rows( |
| raw: torch.Tensor, start: int, end: int, cols: int |
| ) -> torch.Tensor: |
| """Dequant Q4_K rows [start:end]. raw: [out, bytes_per_row] uint8.""" |
| block = 256 |
| elem = 144 |
| chunk = raw[start:end].to(torch.int32) |
| cr = chunk.shape[0] |
| n_blocks = (cols + block - 1) // block |
| chunk = chunk.reshape(cr, n_blocks, elem) |
| d = _f16_view_to_f32(chunk[..., 0:2]) |
| dmin = _f16_view_to_f32(chunk[..., 2:4]) |
| scales = chunk[..., 4:16].to(torch.uint8) |
| qs = chunk[..., 16:144] |
| sc, m = _unpack_scales_k4(scales) |
| sc = sc.to(torch.float32) |
| m = m.to(torch.float32) |
| qs_pairs = qs.reshape(cr, n_blocks, 4, 32) |
| low = (qs_pairs & 0x0F).to(torch.float32) |
| high = (qs_pairs >> 4).to(torch.float32) |
| values = torch.empty(cr, n_blocks, 8, 32, dtype=torch.float32, device=raw.device) |
| values[..., 0::2, :] = low |
| values[..., 1::2, :] = high |
| y = d.unsqueeze(2).unsqueeze(3) * sc.unsqueeze(3) * values \ |
| - dmin.unsqueeze(2).unsqueeze(3) * m.unsqueeze(3) |
| y = y.reshape(cr, n_blocks * block)[:, :cols] |
| return y |
|
|
|
|
| def dequant_q5_k_packed_rows( |
| raw: torch.Tensor, start: int, end: int, cols: int |
| ) -> torch.Tensor: |
| """Dequant Q5_K rows [start:end]. raw: [out, bytes_per_row] uint8.""" |
| block = 256 |
| elem = 176 |
| chunk = raw[start:end].to(torch.int32) |
| cr = chunk.shape[0] |
| n_blocks = (cols + block - 1) // block |
| chunk = chunk.reshape(cr, n_blocks, elem) |
| d = _f16_view_to_f32(chunk[..., 0:2]) |
| dmin = _f16_view_to_f32(chunk[..., 2:4]) |
| scales = chunk[..., 4:16].to(torch.uint8) |
| qh = chunk[..., 16:48] |
| qs = chunk[..., 48:176] |
| sc, m = _unpack_scales_k4(scales) |
| sc = sc.to(torch.float32) |
| m = m.to(torch.float32) |
| qs_pairs = qs.reshape(cr, n_blocks, 4, 32) |
| low = (qs_pairs & 0x0F).to(torch.float32) |
| high = (qs_pairs >> 4).to(torch.float32) |
| values = torch.empty(cr, n_blocks, 8, 32, dtype=torch.float32, device=raw.device) |
| values[..., 0::2, :] = low |
| values[..., 1::2, :] = high |
| bit5 = torch.zeros(cr, n_blocks, 8, 32, dtype=torch.float32, device=raw.device) |
| for k in range(4): |
| bit5[..., 2 * k, :] = ((qh >> (2 * k)) & 1).to(torch.float32) * 16.0 |
| bit5[..., 2 * k + 1, :] = ((qh >> (2 * k + 1)) & 1).to(torch.float32) * 16.0 |
| values5 = values + bit5 |
| y = d.unsqueeze(2).unsqueeze(3) * sc.unsqueeze(3) * values5 \ |
| - dmin.unsqueeze(2).unsqueeze(3) * m.unsqueeze(3) |
| y = y.reshape(cr, n_blocks * block)[:, :cols] |
| return y |
|
|
|
|
| def dequant_q6_k_packed_rows( |
| raw: torch.Tensor, start: int, end: int, cols: int |
| ) -> torch.Tensor: |
| """Dequant Q6_K rows [start:end]. raw: [out, bytes_per_row] uint8. |
| |
| Follows dequantize_row_q6_K (llama.cpp ggml-quants.c): per half-block n, |
| q1/q3 come from ql low/high nibble of first 32 bytes, q2/q4 of bytes 32-64, |
| with scale mapping sc[is+0/1/2/3] (is = l//16) -> sub-blocks 0/4/1/2 |
| interleaved as in the reference. |
| """ |
| block = 256 |
| elem = 210 |
| chunk = raw[start:end].to(torch.int32) |
| cr = chunk.shape[0] |
| n_blocks = (cols + block - 1) // block |
| chunk = chunk.reshape(cr, n_blocks, elem) |
| ql = chunk[..., 0:128] |
| qh = chunk[..., 128:192] |
| sc = chunk[..., 192:208].to(torch.uint8).view(torch.int8).to(torch.float32) |
| d = _f16_view_to_f32(chunk[..., 208:210]) |
| y = torch.zeros(cr, n_blocks, 256, dtype=torch.float32, device=raw.device) |
| for n in range(2): |
| ql_c = ql[:, :, n * 64:(n + 1) * 64] |
| qh_c = qh[:, :, n * 32:(n + 1) * 32] |
| q1 = ((ql_c[:, :, 0:32] & 0x0F) | (((qh_c >> 0) & 3) << 4)).to(torch.float32) - 32.0 |
| q2 = ((ql_c[:, :, 32:64] & 0x0F) | (((qh_c >> 2) & 3) << 4)).to(torch.float32) - 32.0 |
| q3 = ((ql_c[:, :, 0:32] >> 4) | (((qh_c >> 4) & 3) << 4)).to(torch.float32) - 32.0 |
| q4 = ((ql_c[:, :, 32:64] >> 4) | (((qh_c >> 6) & 3) << 4)).to(torch.float32) - 32.0 |
| |
| |
| base = n * 8 |
| d4 = d.unsqueeze(-1) |
| off = n * 128 |
| y[:, :, off + 0:off + 16] = d4 * sc[:, :, base + 0].unsqueeze(-1) * q1[:, :, 0:16] |
| y[:, :, off + 16:off + 32] = d4 * sc[:, :, base + 1].unsqueeze(-1) * q1[:, :, 16:32] |
| y[:, :, off + 32:off + 48] = d4 * sc[:, :, base + 2].unsqueeze(-1) * q2[:, :, 0:16] |
| y[:, :, off + 48:off + 64] = d4 * sc[:, :, base + 3].unsqueeze(-1) * q2[:, :, 16:32] |
| y[:, :, off + 64:off + 80] = d4 * sc[:, :, base + 4].unsqueeze(-1) * q3[:, :, 0:16] |
| y[:, :, off + 80:off + 96] = d4 * sc[:, :, base + 5].unsqueeze(-1) * q3[:, :, 16:32] |
| y[:, :, off + 96:off + 112] = d4 * sc[:, :, base + 6].unsqueeze(-1) * q4[:, :, 0:16] |
| y[:, :, off + 112:off + 128] = d4 * sc[:, :, base + 7].unsqueeze(-1) * q4[:, :, 16:32] |
| return y.reshape(cr, n_blocks * block)[:, :cols] |
|
|
|
|
| |
| _GGUF_SLICE_DEQUANT = { |
| GGML_TYPE_Q8_0: (dequant_q8_0_packed_rows, 32, 34), |
| GGML_TYPE_Q4_K: (dequant_q4_k_packed_rows, 256, 144), |
| GGML_TYPE_Q5_K: (dequant_q5_k_packed_rows, 256, 176), |
| GGML_TYPE_Q6_K: (dequant_q6_k_packed_rows, 256, 210), |
| GGML_TYPE_Q2_K: (dequant_q2_k_packed_rows, 256, 84), |
| GGML_TYPE_Q3_K: (dequant_q3_k_packed_rows, 256, 110), |
| |
| GGML_TYPE_IQ2_XXS: (dequant_iq2_xxs_packed_rows, 256, 66), |
| GGML_TYPE_IQ2_XS: (dequant_iq2_xs_packed, 256, 74), |
| GGML_TYPE_IQ2_S: (dequant_iq2_s_packed, 256, 82), |
| GGML_TYPE_IQ3_XXS: (dequant_iq3_xxs_packed, 256, 98), |
| GGML_TYPE_IQ3_S: (dequant_iq3_s_packed, 256, 110), |
| GGML_TYPE_IQ1_S: (dequant_iq1_s_packed, 256, 66), |
| GGML_TYPE_IQ1_M: (dequant_iq1_m_packed, 256, 56), |
| GGML_TYPE_IQ4_NL: (dequant_iq4_nl_packed, 32, 18), |
| GGML_TYPE_IQ4_XS: (dequant_iq4_xs_packed, 256, 136), |
| } |
|
|
|
|
| def supported_gguf_dtypes() -> set[int]: |
| """GGML dtype ids that support packed slice-dequant.""" |
| return set(_GGUF_SLICE_DEQUANT.keys()) |
|
|
|
|
| def dequant_gguf_slice( |
| weight_raw: torch.Tensor, |
| gguf_dtype: int, |
| start: int, |
| end: int, |
| cols: int, |
| ) -> torch.Tensor: |
| """Dequantize output rows [start:end] from packed GGUF k-quant bytes. |
| |
| Args: |
| weight_raw: uint8 tensor [out_total, bytes_per_row] of packed GGUF data. |
| gguf_dtype: GGML_TYPE_* id (Q8_0, Q4_K, Q5_K, Q6_K, Q2_K, Q3_K). |
| start, end: output row range to dequant (0-based, end exclusive). |
| cols: in_features — trim the dequantized output to this many columns |
| (the last block may be padded beyond cols). |
| |
| Returns: |
| fp32 tensor [end-start, cols]. |
| """ |
| entry = _GGUF_SLICE_DEQUANT.get(gguf_dtype) |
| if entry is None: |
| from agiws_neural_quant.converters.gguf_reader import GGML_TYPE_NAMES |
| raise ValueError( |
| f"dequant_gguf_slice: dtype id={gguf_dtype} " |
| f"({GGML_TYPE_NAMES.get(gguf_dtype, '?')}) not supported. " |
| f"Supported: {sorted(GGML_TYPE_NAMES[d] for d in _GGUF_SLICE_DEQUANT)}" |
| ) |
| fn, _block, _bpb = entry |
| return fn(weight_raw, start, end, cols) |
|
|
|
|
| def bytes_per_row(gguf_dtype: int, cols: int) -> int: |
| """Bytes per output row for a given GGUF dtype and column count.""" |
| entry = _GGUF_SLICE_DEQUANT.get(gguf_dtype) |
| if entry is None: |
| raise ValueError(f"bytes_per_row: unsupported dtype {gguf_dtype}") |
| _fn, block, bpb = entry |
| n_blocks = (cols + block - 1) // block |
| return n_blocks * bpb |
|
|
|
|
| __all__ = [ |
| "dequant_gguf_slice", |
| "dequant_q8_0_packed_rows", |
| "dequant_q4_k_packed_rows", |
| "dequant_q5_k_packed_rows", |
| "dequant_q6_k_packed_rows", |
| "dequant_q2_k_packed_rows", |
| "dequant_q3_k_packed_rows", |
| "dequant_iq2_xxs_packed_rows", |
| "dequant_iq2_xs_packed", |
| "dequant_iq2_s_packed", |
| "dequant_iq3_xxs_packed", |
| "dequant_iq3_s_packed", |
| "dequant_iq1_s_packed", |
| "dequant_iq1_m_packed", |
| "dequant_iq4_nl_packed", |
| "dequant_iq4_xs_packed", |
| "supported_gguf_dtypes", |
| "bytes_per_row", |
| ] |