| """kquant — GGUF k-quant block-wise quantization formats (llama.cpp). |
| |
| Block-wise quantization with super-block structure. Each format packs N weights |
| per block with a shared scale (and optionally min/d-scale). |
| |
| Formats: |
| Q4_0 — block_size=32, fp16 scale, 4-bit values. ~4.5 bpw. |
| Q4_1 — block_size=32, fp16 scale + fp16 min, 4-bit values. ~5 bpw. |
| Q4_K — super-block 256, 8 sub-blocks of 32. 6-bit scale + 6-bit d-scale + 4-bit |
| values. ~4.5 bpw. _S/_M/_L variants differ in d-scale precision. |
| Q5_K — super-block 256, 5-bit values + 6-bit scale + 6-bit d-scale. ~5.5 bpw. |
| Q6_K — super-block 256, 6-bit values + 8-bit d-scale + 6-bit scale. ~6.5 bpw. |
| Q8_0 — block_size=32, fp16 scale, 8-bit values. ~8.5 bpw. Near-lossless. |
| Q2_K — super-block 256, 4-bit Q2-quants + 4-bit d-scale. ~2.6 bpw. |
| Q3_K — super-block 256, 3-bit values + 6-bit scale. ~3.5 bpw. |
| |
| For NeuralQuant we implement the math (quantize/dequantize per block); packing |
| into GGUF binary layout is handled by the GGUF writer (stage 22.8 converters). |
| Here we store block data as tensors and dequantize on-the-fly for inference. |
| """ |
|
|
| from __future__ import annotations |
|
|
| import torch |
|
|
|
|
| |
| |
| |
|
|
| def quantize_q8_0_block(w: torch.Tensor) -> dict[str, torch.Tensor]: |
| """Quantize a block of 32 values to Q8_0. |
| |
| Returns dict with 'scale' (fp32 scalar) and 'values' (int8 [32]). |
| scale = max(abs(w)) / 127; values = round(w / scale). |
| """ |
| assert w.numel() == 32, f"Q8_0 block must be 32 elements, got {w.numel()}" |
| max_abs = w.abs().amax().clamp(min=1e-8) |
| scale = max_abs / 127.0 |
| values = torch.clamp(torch.round(w / scale), min=-127, max=127).to(torch.int8) |
| return {"scale": scale.to(torch.float32), "values": values} |
|
|
|
|
| def dequantize_q8_0_block(scale: torch.Tensor, values: torch.Tensor) -> torch.Tensor: |
| """Reconstruct 32 values from Q8_0 block.""" |
| return values.to(torch.float32) * scale.to(torch.float32) |
|
|
|
|
| |
| |
| |
|
|
| def quantize_q4_0_block(w: torch.Tensor) -> dict[str, torch.Tensor]: |
| assert w.numel() == 32, f"Q4_0 block must be 32 elements, got {w.numel()}" |
| max_abs = w.abs().amax().clamp(min=1e-8) |
| scale = max_abs / 7.0 |
| values = torch.clamp(torch.round(w / scale), min=-8, max=7).to(torch.int8) |
| return {"scale": scale.to(torch.float32), "values": values} |
|
|
|
|
| def dequantize_q4_0_block(scale: torch.Tensor, values: torch.Tensor) -> torch.Tensor: |
| return values.to(torch.float32) * scale.to(torch.float32) |
|
|
|
|
| |
| |
| |
| |
|
|
| def quantize_q4_k_superblock(w: torch.Tensor) -> dict[str, torch.Tensor]: |
| """Quantize 256 values to Q4_K super-block. |
| |
| Structure: |
| - 8 sub-blocks of 32 values, each 4-bit. |
| - super-block scale (fp32, derived from sub-block scales). |
| - per-sub-block d-scale (fp32, ratio sub-scale / super-scale). |
| |
| For NeuralQuant inference we store the actual sub-block scales (not the |
| 6-bit packed GGUF representation); the GGUF writer (stage 22.8) handles |
| bit-packing. |
| """ |
| assert w.numel() == 256, f"Q4_K super-block must be 256 elements, got {w.numel()}" |
| sub = w.reshape(8, 32) |
| |
| sub_scales = sub.abs().amax(dim=1).clamp(min=1e-8) / 7.0 |
| values = torch.clamp(torch.round(sub / sub_scales.unsqueeze(1)), min=-8, max=7).to(torch.int8) |
| |
| super_scale = sub_scales.amax().clamp(min=1e-8) |
| d_scales = sub_scales / super_scale |
| return { |
| "super_scale": super_scale.to(torch.float32), |
| "d_scales": d_scales.to(torch.float32), |
| "values": values.reshape(256), |
| } |
|
|
|
|
| def dequantize_q4_k_superblock(super_scale, d_scales, values) -> torch.Tensor: |
| """Reconstruct 256 values from Q4_K super-block.""" |
| vals = values.to(torch.int8).reshape(8, 32) |
| sub_scales = super_scale.to(torch.float32) * d_scales.to(torch.float32) |
| return (vals.to(torch.float32) * sub_scales.unsqueeze(1)).reshape(256) |
|
|
|
|
| |
| |
| |
|
|
| _Q6_LEVELS = 31 |
|
|
|
|
| def quantize_q6_k_superblock(w: torch.Tensor) -> dict[str, torch.Tensor]: |
| assert w.numel() == 256, f"Q6_K super-block must be 256 elements, got {w.numel()}" |
| sub = w.reshape(8, 32) |
| sub_scales = sub.abs().amax(dim=1).clamp(min=1e-8) / _Q6_LEVELS |
| values = torch.clamp(torch.round(sub / sub_scales.unsqueeze(1)), min=-32, max=31).to(torch.int8) |
| super_scale = sub_scales.amax().clamp(min=1e-8) |
| d_scales = sub_scales / super_scale |
| return { |
| "super_scale": super_scale.to(torch.float32), |
| "d_scales": d_scales.to(torch.float32), |
| "values": values.reshape(256), |
| } |
|
|
|
|
| def dequantize_q6_k_superblock(super_scale, d_scales, values) -> torch.Tensor: |
| vals = values.to(torch.int8).reshape(8, 32) |
| sub_scales = super_scale.to(torch.float32) * d_scales.to(torch.float32) |
| return (vals.to(torch.float32) * sub_scales.unsqueeze(1)).reshape(256) |
|
|
|
|
| |
| |
| |
|
|
| _Q5_LEVELS = 15 |
|
|
|
|
| def quantize_q5_k_superblock(w: torch.Tensor) -> dict[str, torch.Tensor]: |
| assert w.numel() == 256 |
| sub = w.reshape(8, 32) |
| sub_scales = sub.abs().amax(dim=1).clamp(min=1e-8) / _Q5_LEVELS |
| values = torch.clamp(torch.round(sub / sub_scales.unsqueeze(1)), min=-16, max=15).to(torch.int8) |
| super_scale = sub_scales.amax().clamp(min=1e-8) |
| d_scales = sub_scales / super_scale |
| return { |
| "super_scale": super_scale.to(torch.float32), |
| "d_scales": d_scales.to(torch.float32), |
| "values": values.reshape(256), |
| } |
|
|
|
|
| def dequantize_q5_k_superblock(super_scale, d_scales, values) -> torch.Tensor: |
| vals = values.to(torch.int8).reshape(8, 32) |
| sub_scales = super_scale.to(torch.float32) * d_scales.to(torch.float32) |
| return (vals.to(torch.float32) * sub_scales.unsqueeze(1)).reshape(256) |
|
|
|
|
| |
| |
| |
|
|
| _Q2_LEVELS = 1 |
|
|
|
|
| def quantize_q2_k_superblock(w: torch.Tensor) -> dict[str, torch.Tensor]: |
| assert w.numel() == 256 |
| sub = w.reshape(8, 32) |
| sub_scales = sub.abs().amax(dim=1).clamp(min=1e-8) / _Q2_LEVELS |
| values = torch.clamp(torch.round(sub / sub_scales.unsqueeze(1)), min=-2, max=1).to(torch.int8) |
| super_scale = sub_scales.amax().clamp(min=1e-8) |
| d_scales = sub_scales / super_scale |
| return { |
| "super_scale": super_scale.to(torch.float32), |
| "d_scales": d_scales.to(torch.float32), |
| "values": values.reshape(256), |
| } |
|
|
|
|
| def dequantize_q2_k_superblock(super_scale, d_scales, values) -> torch.Tensor: |
| vals = values.to(torch.int8).reshape(8, 32) |
| sub_scales = super_scale.to(torch.float32) * d_scales.to(torch.float32) |
| return (vals.to(torch.float32) * sub_scales.unsqueeze(1)).reshape(256) |
|
|
|
|
| |
| |
| |
|
|
| _Q3_LEVELS = 3 |
|
|
|
|
| def quantize_q3_k_superblock(w: torch.Tensor) -> dict[str, torch.Tensor]: |
| assert w.numel() == 256 |
| sub = w.reshape(8, 32) |
| sub_scales = sub.abs().amax(dim=1).clamp(min=1e-8) / _Q3_LEVELS |
| values = torch.clamp(torch.round(sub / sub_scales.unsqueeze(1)), min=-4, max=3).to(torch.int8) |
| super_scale = sub_scales.amax().clamp(min=1e-8) |
| d_scales = sub_scales / super_scale |
| return { |
| "super_scale": super_scale.to(torch.float32), |
| "d_scales": d_scales.to(torch.float32), |
| "values": values.reshape(256), |
| } |
|
|
|
|
| def dequantize_q3_k_superblock(super_scale, d_scales, values) -> torch.Tensor: |
| vals = values.to(torch.int8).reshape(8, 32) |
| sub_scales = super_scale.to(torch.float32) * d_scales.to(torch.float32) |
| return (vals.to(torch.float32) * sub_scales.unsqueeze(1)).reshape(256) |
|
|
|
|
| |
| |
| |
|
|
| BLOCK_SIZE_32 = 32 |
| BLOCK_SIZE_256 = 256 |
|
|
| FORMATS = { |
| "q8_0": (quantize_q8_0_block, dequantize_q8_0_block, BLOCK_SIZE_32), |
| "q4_0": (quantize_q4_0_block, dequantize_q4_0_block, BLOCK_SIZE_32), |
| "q4_k": (quantize_q4_k_superblock, dequantize_q4_k_superblock, BLOCK_SIZE_256), |
| "q5_k": (quantize_q5_k_superblock, dequantize_q5_k_superblock, BLOCK_SIZE_256), |
| "q6_k": (quantize_q6_k_superblock, dequantize_q6_k_superblock, BLOCK_SIZE_256), |
| "q2_k": (quantize_q2_k_superblock, dequantize_q2_k_superblock, BLOCK_SIZE_256), |
| "q3_k": (quantize_q3_k_superblock, dequantize_q3_k_superblock, BLOCK_SIZE_256), |
| } |
|
|
|
|
| def quantize_blocks(w: torch.Tensor, fmt: str) -> dict[str, torch.Tensor]: |
| """Quantize a flat weight tensor into blocks of the given format. |
| |
| Pads the last dim to be divisible by block_size. |
| |
| Returns dict with stacked block tensors. |
| """ |
| quant_fn, _, block_size = FORMATS[fmt] |
| out_features = w.shape[0] |
| in_features = w.shape[1] if w.dim() > 1 else w.numel() |
| if w.dim() > 1: |
| flat = w |
| else: |
| flat = w.reshape(1, -1) |
| out_features = 1 |
| |
| pad = (block_size - (flat.shape[1] % block_size)) % block_size |
| if pad > 0: |
| flat = torch.nn.functional.pad(flat, (0, pad)) |
| in_padded = flat.shape[1] |
| num_blocks = in_padded // block_size |
|
|
| |
| blocks = flat.reshape(out_features, num_blocks, block_size) |
|
|
| if fmt in ("q8_0", "q4_0"): |
| |
| n_levels = 127 if fmt == "q8_0" else 7 |
| max_val = n_levels if fmt == "q8_0" else 7 |
| min_val = -n_levels if fmt == "q8_0" else -8 |
| scales = blocks.abs().amax(dim=2).clamp(min=1e-8) / n_levels |
| values = torch.clamp( |
| torch.round(blocks / scales.unsqueeze(2)), min=min_val, max=max_val |
| ).to(torch.int8) |
| return { |
| "scales": scales.to(torch.float32), |
| "values": values, |
| "in_features": in_features, |
| "in_padded": in_padded, |
| "out_features": out_features, |
| "block_size": block_size, |
| } |
| else: |
| |
| |
| n_levels_map = {"q4_k": 7, "q5_k": 15, "q6_k": 31, "q2_k": 1, "q3_k": 3} |
| min_val_map = {"q4_k": -8, "q5_k": -16, "q6_k": -32, "q2_k": -2, "q3_k": -4} |
| n_levels = n_levels_map[fmt] |
| min_val = min_val_map[fmt] |
| sub = blocks.reshape(out_features, num_blocks, 8, 32) |
| sub_scales = sub.abs().amax(dim=3).clamp(min=1e-8) / n_levels |
| values = torch.clamp( |
| torch.round(sub / sub_scales.unsqueeze(3)), min=min_val, max=n_levels |
| ).to(torch.int8) |
| super_scales = sub_scales.amax(dim=2).clamp(min=1e-8) |
| d_scales = sub_scales / super_scales.unsqueeze(2) |
| return { |
| "super_scales": super_scales.to(torch.float32), |
| "d_scales": d_scales.to(torch.float32), |
| "values": values.reshape(out_features, num_blocks, 256), |
| "in_features": in_features, |
| "in_padded": in_padded, |
| "out_features": out_features, |
| "block_size": 256, |
| "num_super": num_blocks, |
| } |
|
|
|
|
| def dequantize_blocks(qd: dict[str, torch.Tensor], fmt: str) -> torch.Tensor: |
| """Reconstruct the weight tensor from block-quantized data.""" |
| in_features = qd["in_features"] |
| in_padded = qd["in_padded"] |
| out_features = qd["out_features"] |
| if fmt in ("q8_0", "q4_0"): |
| scales = qd["scales"] |
| values = qd["values"] |
| w = values.to(torch.float32) * scales.to(torch.float32).unsqueeze(2) |
| w = w.reshape(out_features, in_padded)[:, :in_features] |
| else: |
| super_scales = qd["super_scales"] |
| d_scales = qd["d_scales"] |
| values = qd["values"] |
| sub_scales = super_scales.to(torch.float32).unsqueeze(2) * d_scales.to(torch.float32) |
| sub = values.reshape(out_features, -1, 8, 32) |
| w = sub.to(torch.float32) * sub_scales.unsqueeze(3) |
| w = w.reshape(out_features, -1)[:, :in_features] |
| return w |