Spaces:
Sleeping
Sleeping
File size: 8,348 Bytes
28a1a01 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 174 175 176 177 178 179 180 181 182 183 184 185 186 187 188 189 190 191 192 193 194 195 196 197 198 199 200 201 | """4-bit weight quantization study: element format x scale precision x extras.
CISM's ``hybrid-int4`` packs MLP weights as symmetric INT4 (levels -7..7,
blocks of 32, per-block FP32 scale = amax/7, round-half-away-from-zero) with
everything else 2D per-row INT8. This study holds that pipeline fixed and
swaps one knob at a time:
* element format: uniform INT4 vs E2M1 (MXFP4) vs NF4 (QLoRA)
* scale precision: FP32 vs FP8 E4M3 vs power-of-two E8M0 (OCP MXFP4)
* block size: 32 vs 16 (NVFP4 granularity)
* error feedback: per-block residual diffusion on top of the winner
The NumPy simulation mirrors the native quantizer bit-for-bit (validated
against the real engine below), rebuilds FP32 engines from the dequantized
weights, and scores perplexity on the AutoBench corpus over small and not-so-
small checkpoints. Quality only; speed is a kernel question, decided later.
"""
import argparse
import json
import numpy as np
import torch
from cism import Engine
from cism.autobench import CORPUS
from cism.loader import import_model
E2M1 = np.array([0.0, 0.5, 1.0, 1.5, 2.0, 3.0, 4.0, 6.0], dtype=np.float64)
NF4 = np.array([-1.0, -0.6962, -0.5239, -0.3650, -0.1848, -0.1274, -0.0400, 0.0,
0.0397, 0.1292, 0.2826, 0.5626, 0.5824, 0.8400, 1.0770, 1.2712])
TINY = float(np.finfo(np.float32).tiny)
def round_half_away(values: np.ndarray) -> np.ndarray:
# std::round semantics: halfway cases go away from zero. np.round does
# banker's rounding and would NOT mirror the native quantizer.
return np.floor(np.abs(values) + 0.5) * np.sign(values)
def _blocked(matrix: np.ndarray, block: int) -> np.ndarray:
rows, cols = matrix.shape
pad = (-cols) % block
padded = np.pad(matrix, ((0, 0), (0, pad)))
return padded.reshape(rows, -1, block)
def _scale_for(amax: np.ndarray, mode: str, qmax: float) -> np.ndarray:
if mode == "fp32":
return np.maximum(amax / qmax, TINY)
if mode == "e8m0":
# OCP shared exponent: power of two, floor(log2(amax)) - (2 or 3) so
# the largest element value can reach but not exceed amax.
return np.exp2(np.floor(np.log2(np.maximum(amax, TINY))) - np.log2(qmax))
if mode == "e4m3":
base = torch.from_numpy(np.maximum(amax / qmax, TINY))
return np.maximum(base.to(torch.float8_e4m3fn).to(torch.float32).numpy(), TINY)
raise ValueError(scale_mode := mode)
def _lut_round(magnitude: np.ndarray, lut: np.ndarray) -> np.ndarray:
index = np.argmin(np.abs(magnitude[..., None] - lut), axis=-1)
return lut[index]
def _deblock(dequantized: np.ndarray, rows: int, cols: int) -> np.ndarray:
return dequantized.reshape(rows, -1)[:, :cols].astype(np.float32)
def quantize_uniform(matrix: np.ndarray, *, levels: int = 7, scale_mode: str = "fp32",
block: int = 32) -> np.ndarray:
"""Symmetric integer grid (-levels..levels), mirrors native int4."""
rows, cols = matrix.shape
blocks = _blocked(matrix, block)
amax = np.abs(blocks).max(axis=-1, keepdims=True)
scale = _scale_for(amax, scale_mode, float(levels))
quantized = np.clip(round_half_away(blocks / scale), -levels, levels)
return _deblock(quantized * scale, rows, cols)
def quantize_fp_e2m1(matrix: np.ndarray, *, scale_mode: str = "e4m3",
block: int = 32) -> np.ndarray:
rows, cols = matrix.shape
blocks = _blocked(matrix, block)
amax = np.abs(blocks).max(axis=-1, keepdims=True)
scale = _scale_for(amax, scale_mode, 6.0)
dequantized = np.sign(blocks) * _lut_round(np.abs(blocks) / scale, E2M1) * scale
return _deblock(dequantized, rows, cols)
def quantize_nf4(matrix: np.ndarray, *, scale_mode: str = "fp32",
block: int = 32) -> np.ndarray:
rows, cols = matrix.shape
blocks = _blocked(matrix, block)
amax = np.abs(blocks).max(axis=-1, keepdims=True)
scale = _scale_for(amax, scale_mode, 1.0)
dequantized = _lut_round(blocks / scale, NF4) * scale
return _deblock(dequantized, rows, cols)
def quantize_uniform_feedback(matrix: np.ndarray, *, levels: int = 7) -> np.ndarray:
"""Fixed-scale RTN plus per-block residual diffusion along the row."""
rows, cols = matrix.shape
blocks = _blocked(matrix, 32)
amax = np.abs(blocks).max(axis=-1)
scale = np.maximum(amax / float(levels), TINY) # (rows, nblocks)
quantized = np.zeros_like(blocks)
residual = np.zeros((rows, blocks.shape[1]))
for position in range(blocks.shape[-1]):
target = blocks[:, :, position] + residual
quantized[:, :, position] = np.clip(round_half_away(target / scale), -levels, levels)
residual = target - quantized[:, :, position] * scale
return _deblock(quantized * scale[..., None], rows, cols)
def quantize_e2m1_feedback(matrix: np.ndarray, *, scale_mode: str = "e4m3") -> np.ndarray:
rows, cols = matrix.shape
blocks = _blocked(matrix, 32)
amax = np.abs(blocks).max(axis=-1)
scale = _scale_for(amax, scale_mode, 6.0) # (rows, nblocks)
quantized = np.zeros_like(blocks)
residual = np.zeros((rows, blocks.shape[1]))
for position in range(blocks.shape[-1]):
target = blocks[:, :, position] + residual
quantized[:, :, position] = np.sign(target) * _lut_round(np.abs(target) / scale, E2M1)
residual = target - quantized[:, :, position] * scale
return _deblock(quantized * scale[..., None], rows, cols)
def quantize_int8_row(matrix: np.ndarray) -> np.ndarray:
"""Mirror of the native Storage::int8 constructor (per-row scale)."""
amax = np.abs(matrix).max(axis=1, keepdims=True)
scale = np.maximum(amax / 127.0, TINY)
quantized = np.clip(round_half_away(matrix / scale), -127, 127)
return (quantized * scale).astype(np.float32)
def simulate(weights: dict, scheme) -> dict:
out = {}
for name, array in weights.items():
if array.ndim == 1:
out[name] = array.astype(np.float32)
elif ".mlp." in name:
out[name] = scheme(array)
else:
out[name] = quantize_int8_row(array)
return out
def perplexity(config: dict, weights: dict, tokenizer) -> float:
engine = Engine.from_weights(config, weights, tokenizer, model_id="study")
return engine.nll(CORPUS, window=128)["perplexity"]
SCHEMES = {
"int4 fp32 b32 (ours)": lambda m: quantize_uniform(m),
"int4 e4m3 b32": lambda m: quantize_uniform(m, scale_mode="e4m3"),
"int4 e8m0 b32": lambda m: quantize_uniform(m, scale_mode="e8m0"),
"E2M1 e8m0 b32 (OCP MXFP4)": lambda m: quantize_fp_e2m1(m, scale_mode="e8m0"),
"E2M1 e4m3 b32 (NVFP4-ish)": lambda m: quantize_fp_e2m1(m, scale_mode="e4m3"),
"E2M1 e4m3 b16": lambda m: quantize_fp_e2m1(m, scale_mode="e4m3", block=16),
"int4 e4m3 b16": lambda m: quantize_uniform(m, scale_mode="e4m3", block=16),
"E2M1 fp32 b32": lambda m: quantize_fp_e2m1(m, scale_mode="fp32"),
"NF4 fp32 b32 (QLoRA)": lambda m: quantize_nf4(m, scale_mode="fp32"),
"NF4 e4m3 b32": lambda m: quantize_nf4(m, scale_mode="e4m3"),
"int4 fb fp32 b32": lambda m: quantize_uniform_feedback(m),
"E2M1 fb e4m3 b32": lambda m: quantize_e2m1_feedback(m),
}
def main() -> int:
parser = argparse.ArgumentParser()
parser.add_argument("models", nargs="+")
args = parser.parse_args()
results = {}
for model_id in args.models:
loaded = import_model(model_id, local_files_only=True)
config, weights, tokenizer = dict(loaded.config), loaded.weights, loaded.tokenizer
row = {}
row["fp32 (reference)"] = perplexity(config, weights, tokenizer)
row["hybrid-int4 (native engine)"] = Engine.from_pretrained(
model_id, precision="hybrid-int4", local_files_only=True).nll(CORPUS, window=128)["perplexity"]
for label, scheme in SCHEMES.items():
row[label] = perplexity(config, simulate(weights, scheme), tokenizer)
results[model_id] = row
reference = row["fp32 (reference)"]
print(f"\n=== {model_id} (fp32 PPL {reference:.2f}) ===")
for label, value in row.items():
delta = 100 * (value - reference) / reference
print(f" {label:32s} PPL {value:10.3f} ({delta:+.3f}%)")
print("\n" + json.dumps(results, indent=2))
return 0
if __name__ == "__main__":
raise SystemExit(main())
|