code/10-finetune-deploy/quantize.py
89 lines · 3.5 KBCode and program output are shown exactly as they ran, so comments and printed output are in Chinese.
"""手写量化:把第 09 模块的小 GPT 的权重从 32 位小数压成 8 比特、4 比特整数,看看它变小了多少、变差了多少。
python quantize.py
先在 code/09-transformer 下运行 train.py,得到 .cache/gpt.pt。只用 CPU,十几秒。
"""
import sys
from pathlib import Path
import torch
GPT_DIR = Path(__file__).parent.parent / "09-transformer"
sys.path.insert(0, str(GPT_DIR))
from gpt import GPT, Config # noqa: E402
ckpt = torch.load(GPT_DIR / ".cache" / "gpt.pt")
chars = ckpt["chars"]
index = {c: i for i, c in enumerate(chars)}
def quantize(w, bits, group=None):
"""对称量化:每组数用一个缩放系数,把 [-最大绝对值, 最大绝对值] 映射到整数 [-qmax, qmax]。
返回"量化后再还原"的权重,以及实际要存的整数和缩放系数。"""
qmax = 2 ** (bits - 1) - 1 # 8 比特是 127,4 比特是 7
shape = w.shape
w = w.reshape(-1, group) if group else w.reshape(shape[0], -1) # 按组,或者按行
scale = w.abs().amax(dim=1, keepdim=True) / qmax
q = torch.round(w / scale).clamp(-qmax, qmax) # 这就是要存下来的整数
return (q * scale).reshape(shape), q, scale
print("== 1. 一个例子:把 8 个小数量化成 4 比特整数")
torch.manual_seed(0)
w = torch.randn(1, 8) * 0.05
deq, q, scale = quantize(w, 4)
print(f" 原来: {[round(x, 4) for x in w[0].tolist()]}")
print(f" 整数: {[int(x) for x in q[0].tolist()]}(缩放系数 {scale.item():.5f})")
print(f" 还原后: {[round(x, 4) for x in deq[0].tolist()]}")
def load(bits=None, group=None):
model = GPT(Config(**ckpt["config"]))
model.load_state_dict(ckpt["model"])
if bits:
with torch.no_grad():
for name, p in model.named_parameters():
if p.dim() == 2: # 只量化矩阵(线性层和嵌入),LayerNorm 和偏置很小,保持原样
p.copy_(quantize(p, bits, group)[0])
return model.eval()
def size_mb(model, bits=None, group=None):
total = 0
for p in model.parameters():
if bits and p.dim() == 2:
n_scales = p.numel() // group if group else p.shape[0]
total += p.numel() * bits / 8 + n_scales * 2 # 整数 + 每组一个 16 位的缩放系数
else:
total += p.numel() * 4
return total / 1024 / 1024
poems = (GPT_DIR / ".cache" / "poems.txt").read_text().splitlines()
val = torch.tensor([index[c] for c in "\n".join(poems[-1000:]) + "\n"])
@torch.no_grad()
def val_loss(model):
g = torch.Generator().manual_seed(123)
starts = torch.randint(len(val) - 129, (200,), generator=g)
x = torch.stack([val[s:s + 128] for s in starts])
y = torch.stack([val[s + 1:s + 129] for s in starts])
return model(x, y)[1].item()
def poem(model):
torch.manual_seed(0)
ids = model.generate(torch.tensor([[index["\n"]]]), 80, temperature=0.8, stop_id=index["\n"])
return "".join(chars[i] for i in ids[0, 1:].tolist()).strip()
print("\n== 2. 整个模型量化之后")
base = load()
print(f" {'':<24}{'大小':>8}{'验证损失':>10}")
for label, bits, group in [("32 位小数(原模型)", None, None), ("8 比特,每行一个系数", 8, None),
("4 比特,每行一个系数", 4, None), ("4 比特,每 32 个数一个系数", 4, 32),
("2 比特,每 32 个数一个系数", 2, 32)]:
m = load(bits, group)
print(f" {label:<20}{size_mb(base, bits, group):8.2f} MB{val_loss(m):10.3f} {poem(m)[:24]}")