"""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 # --------------------------------------------------------------------------- # Q8_0 — block_size=32, fp16 scale, int8 values. Near-lossless int8. # --------------------------------------------------------------------------- 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) # --------------------------------------------------------------------------- # Q4_0 — block_size=32, fp16 scale, 4-bit values [-8, 7]. # --------------------------------------------------------------------------- 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 # 4-bit symmetric: [-8, 7], use 7 for scale 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) # --------------------------------------------------------------------------- # Q4_K — super-block 256 = 8 sub-blocks of 32. 6-bit packed scale + 6-bit # d-scale per sub-block + 4-bit values. # --------------------------------------------------------------------------- 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) # Per-sub-block scale: absmax / 7 (4-bit symmetric). sub_scales = sub.abs().amax(dim=1).clamp(min=1e-8) / 7.0 # [8] values = torch.clamp(torch.round(sub / sub_scales.unsqueeze(1)), min=-8, max=7).to(torch.int8) # Super-block scale = max(sub_scales). d-scale = sub_scale / super_scale. super_scale = sub_scales.amax().clamp(min=1e-8) d_scales = sub_scales / super_scale # [8], in (0, 1] return { "super_scale": super_scale.to(torch.float32), "d_scales": d_scales.to(torch.float32), "values": values.reshape(256), # int8 [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) # [8] return (vals.to(torch.float32) * sub_scales.unsqueeze(1)).reshape(256) # --------------------------------------------------------------------------- # Q6_K — super-block 256, 6-bit values + 8-bit d-scale + 6-bit super-scale. # --------------------------------------------------------------------------- _Q6_LEVELS = 31 # 6-bit symmetric [-32, 31], use 31 for scale 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 # [8] 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 # [8] 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_K — super-block 256, 5-bit values + 6-bit d-scale. # --------------------------------------------------------------------------- _Q5_LEVELS = 15 # 5-bit symmetric [-16, 15], use 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_K — super-block 256, 2-bit values + 4-bit d-scale. Extreme compression. # --------------------------------------------------------------------------- _Q2_LEVELS = 1 # 2-bit symmetric [-2, 1], use 1 for scale (coarse) 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_K — super-block 256, 3-bit values + 6-bit d-scale. # --------------------------------------------------------------------------- _Q3_LEVELS = 3 # 3-bit symmetric [-4, 3], use 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) # --------------------------------------------------------------------------- # Format registry: format string -> (block quantize, block dequantize, block_size) # --------------------------------------------------------------------------- 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 in_features to be divisible by block_size. 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 # Reshape to [out, num_blocks, block_size]. blocks = flat.reshape(out_features, num_blocks, block_size) if fmt in ("q8_0", "q4_0"): # Simple block: per-block absmax scale + int values. 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 # [out, num_blocks] 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: # K-formats: super-block 256 = 8 sub-blocks of 32. # blocks already [out, num_blocks, 256]; reshape to [out, num_blocks, 8, 32]. 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 # [out, num_blocks, 8] 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) # [out, num_blocks] d_scales = sub_scales / super_scales.unsqueeze(2) # [out, num_blocks, 8] 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"] # [out, num_blocks] values = qd["values"] # [out, num_blocks, block_size] 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"] # [out, num_blocks] d_scales = qd["d_scales"] # [out, num_blocks, 8] values = qd["values"] # [out, num_blocks, 256] 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