code/10-finetune-deploy/quantize.py

89 lines · 3.5 KB

Code 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]}")