ArGrigorov's picture
Upload folder using huggingface_hub
e9c8366 verified
Raw
History Blame Contribute Delete
14 kB
"""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