kingjones777's picture
Add files using upload-large-folder tool
da1a4ff verified
Raw History Blame Contribute Delete
27.6 kB
"""CPU tests for weight-only INT8 Ming MLLM quantize + load.
Run: HIP_VISIBLE_DEVICES=-1 python test_int8.py
"""
from __future__ import annotations
import json
import sys
import tempfile
import traceback
from pathlib import Path
import torch
import torch.nn.functional as F
from safetensors.torch import load_file, save_file
from torch import nn
import quantize_stream
from int8_linear import Int8Linear, is_quantizable, quantize_weight
from load_int8 import load_int8_mllm_
# Tiny stand-in for Ming's MLLM names. Not the real model.
HIDDEN = 32
INTER = 48
VOCAB = 64
N_EXPERTS = 2
class RMSNorm(nn.Module):
def __init__(self, dim: int, eps: float = 1e-6):
super().__init__()
self.weight = nn.Parameter(torch.ones(dim))
self.eps = eps
def forward(self, x: torch.Tensor) -> torch.Tensor:
var = x.float().pow(2).mean(dim=-1, keepdim=True)
y = x * torch.rsqrt(var + self.eps)
return (y * self.weight).to(dtype=x.dtype)
class Attention(nn.Module):
def __init__(self, hidden: int):
super().__init__()
self.hidden = hidden
self.query_key_value = nn.Linear(hidden, hidden * 3, bias=True)
self.dense = nn.Linear(hidden, hidden, bias=False)
self.q_norm = RMSNorm(hidden)
self.k_norm = RMSNorm(hidden)
# Non-persistent, like BailingMoeV2RotaryEmbedding.inv_freq.
self.register_buffer(
"inv_freq", torch.arange(hidden // 2, dtype=torch.float32), persistent=False
)
def forward(self, x: torch.Tensor) -> torch.Tensor:
qkv = self.query_key_value(x)
h = self.hidden
q = self.q_norm(qkv[..., :h])
k = self.k_norm(qkv[..., h : 2 * h])
v = qkv[..., 2 * h :]
return self.dense(q + k + v)
class DenseMLP(nn.Module):
def __init__(self, hidden: int, inter: int):
super().__init__()
self.gate_proj = nn.Linear(hidden, inter, bias=False)
self.up_proj = nn.Linear(hidden, inter, bias=True)
self.down_proj = nn.Linear(inter, hidden, bias=False)
def forward(self, x: torch.Tensor) -> torch.Tensor:
return self.down_proj(F.silu(self.gate_proj(x)) * self.up_proj(x))
class Expert(nn.Module):
def __init__(self, hidden: int, inter: int):
super().__init__()
self.gate_proj = nn.Linear(hidden, inter, bias=False)
self.up_proj = nn.Linear(hidden, inter, bias=True)
self.down_proj = nn.Linear(inter, hidden, bias=False)
def forward(self, x: torch.Tensor) -> torch.Tensor:
return self.down_proj(F.silu(self.gate_proj(x)) * self.up_proj(x))
class Router(nn.Module):
"""Not an nn.Linear. Leaf name is gate / image_gate / audio_gate."""
def __init__(self, hidden: int, n_experts: int):
super().__init__()
self.weight = nn.Parameter(torch.empty(n_experts, hidden))
self.expert_bias = nn.Parameter(torch.zeros(n_experts), requires_grad=False)
nn.init.kaiming_uniform_(self.weight, a=5**0.5)
def forward(self, x: torch.Tensor) -> torch.Tensor:
return F.linear(x, self.weight, self.expert_bias)
class MoeMLP(nn.Module):
def __init__(self, hidden: int, inter: int, n_experts: int):
super().__init__()
self.gate = Router(hidden, n_experts)
self.image_gate = Router(hidden, n_experts)
self.audio_gate = Router(hidden, n_experts)
self.experts = nn.ModuleList(Expert(hidden, inter) for _ in range(n_experts))
self.shared_experts = Expert(hidden, inter)
def forward(self, x: torch.Tensor) -> torch.Tensor:
scores = self.gate(x) + self.image_gate(x) + self.audio_gate(x)
weights = torch.softmax(scores, dim=-1)
mixed = self.shared_experts(x)
for i, expert in enumerate(self.experts):
mixed = mixed + expert(x) * weights[..., i : i + 1]
return mixed
class DecoderLayer(nn.Module):
def __init__(self, hidden: int, mlp: nn.Module):
super().__init__()
self.input_layernorm = RMSNorm(hidden)
self.post_attention_layernorm = RMSNorm(hidden)
self.attention = Attention(hidden)
self.mlp = mlp
def forward(self, x: torch.Tensor) -> torch.Tensor:
x = x + self.attention(self.input_layernorm(x))
x = x + self.mlp(self.post_attention_layernorm(x))
return x
class TinyMing(nn.Module):
"""Names match the real checkpoint: model.model.layers.*, model.lm_head, vision.*."""
def __init__(self):
super().__init__()
self.model = nn.Module()
self.model.model = nn.Module()
self.model.model.word_embeddings = nn.Embedding(VOCAB, HIDDEN)
self.model.model.layers = nn.ModuleList(
[
DecoderLayer(HIDDEN, DenseMLP(HIDDEN, INTER)),
DecoderLayer(HIDDEN, MoeMLP(HIDDEN, INTER, N_EXPERTS)),
]
)
self.model.model.norm = RMSNorm(HIDDEN)
self.model.lm_head = nn.Linear(HIDDEN, VOCAB, bias=False)
block = nn.Module()
block.attn = nn.Module()
block.attn.qkv = nn.Linear(HIDDEN, HIDDEN, bias=False)
self.vision = nn.Module()
self.vision.blocks = nn.ModuleList([block])
self.linear_proj = nn.ModuleList([nn.Linear(HIDDEN, HIDDEN, bias=True)])
def forward(self, input_ids: torch.Tensor) -> torch.Tensor:
h = self.model.model.word_embeddings(input_ids)
for layer in self.model.model.layers:
h = layer(h)
h = self.model.model.norm(h)
return self.model.lm_head(h)
# Modules the rule must select for TinyMing. Hardcoded — not derived from is_quantizable.
EXPECTED_QUANT_MODULES = [
"model.model.layers.0.attention.dense",
"model.model.layers.0.attention.query_key_value",
"model.model.layers.0.mlp.down_proj",
"model.model.layers.0.mlp.gate_proj",
"model.model.layers.0.mlp.up_proj",
"model.model.layers.1.attention.dense",
"model.model.layers.1.attention.query_key_value",
"model.model.layers.1.mlp.experts.0.down_proj",
"model.model.layers.1.mlp.experts.0.gate_proj",
"model.model.layers.1.mlp.experts.0.up_proj",
"model.model.layers.1.mlp.experts.1.down_proj",
"model.model.layers.1.mlp.experts.1.gate_proj",
"model.model.layers.1.mlp.experts.1.up_proj",
"model.model.layers.1.mlp.shared_experts.down_proj",
"model.model.layers.1.mlp.shared_experts.gate_proj",
"model.model.layers.1.mlp.shared_experts.up_proj",
]
MUST_NOT_QUANTIZE = [
"model.model.layers.1.mlp.gate",
"model.model.layers.1.mlp.image_gate",
"model.model.layers.1.mlp.audio_gate",
"model.model.word_embeddings",
"model.model.norm",
"model.lm_head",
"vision.blocks.0.attn.qkv",
"linear_proj.0",
"model.model.layers.0.attention.q_norm",
"model.model.layers.0.input_layernorm",
]
def _move_parameters_to_meta(model: nn.Module) -> nn.Module:
"""Parameters → meta, buffers stay where they are (CPU). Matches accelerate include_buffers=False."""
for mod in model.modules():
for name, param in list(mod._parameters.items()):
if param is None:
continue
mod._parameters[name] = nn.Parameter(
param.detach().to(device="meta"),
requires_grad=param.requires_grad,
)
return model
def _save_bf16_checkpoint(model: nn.Module, src: Path) -> None:
src.mkdir(parents=True, exist_ok=True)
sd = {k: v.detach().contiguous() for k, v in model.state_dict().items()}
if not sd:
raise AssertionError("empty state_dict")
for tensor in sd.values():
if tensor.is_floating_point():
assert tensor.dtype == torch.bfloat16, tensor.dtype
keys = list(sd)
mid = max(1, len(keys) // 2)
shards = {
"bf16-00001.safetensors": {k: sd[k] for k in keys[:mid]},
"bf16-00002.safetensors": {k: sd[k] for k in keys[mid:]},
}
weight_map = {}
total = 0
for filename, tensors in shards.items():
save_file(tensors, str(src / filename))
for name, tensor in tensors.items():
weight_map[name] = filename
total += tensor.numel() * tensor.element_size()
index = {"metadata": {"total_size": total}, "weight_map": weight_map}
(src / "model.safetensors.index.json").write_text(
json.dumps(index, indent=2) + "\n", encoding="utf-8"
)
(src / "config.json").write_bytes(b'{"model_type":"tiny-ming","hidden":32}\n')
extra = src / "extra"
extra.mkdir()
(extra / "chat_template.jinja").write_text("{{ messages }}\n", encoding="utf-8")
def _load_all(folder: Path) -> dict[str, torch.Tensor]:
index = json.loads((folder / "model.safetensors.index.json").read_text(encoding="utf-8"))
order: list[str] = []
seen: set[str] = set()
for shard in index["weight_map"].values():
if shard not in seen:
seen.add(shard)
order.append(shard)
sd: dict[str, torch.Tensor] = {}
for shard in order:
sd.update(load_file(str(folder / shard)))
return sd
def _apply_int8_(model: nn.Module) -> None:
names = []
for name, mod in model.named_modules():
if isinstance(mod, nn.Linear) and is_quantizable(
f"{name}.weight", tuple(mod.weight.shape)
):
names.append(name)
for name in names:
parent_name, _, leaf = name.rpartition(".")
parent = model.get_submodule(parent_name) if parent_name else model
setattr(parent, leaf, Int8Linear.from_linear(getattr(parent, leaf)))
def _assert_no_meta(model: nn.Module) -> None:
for name, param in model.named_parameters():
assert param.device.type != "meta", name
for mod_name, mod in model.named_modules():
for buf_name, buf in mod._buffers.items():
if buf is None:
continue
full = f"{mod_name}.{buf_name}" if mod_name else buf_name
assert buf.device.type != "meta", full
def test_from_linear_roundtrip() -> None:
torch.manual_seed(0)
out_f, in_f = 5, 7
lin = nn.Linear(in_f, out_f, bias=True)
scales = torch.tensor([0.5, 0.25, 0.125, 2.0, 4.0], dtype=torch.float32)
q = torch.randint(-127, 128, (out_f, in_f), dtype=torch.int8)
q[:, 0] = 127
q[2, :] = 0 # all-zero row; must not NaN
weight = q.float() * scales[:, None]
with torch.no_grad():
lin.weight.copy_(weight)
lin.bias.copy_(torch.tensor([0.1, -0.2, 0.3, -0.4, 0.5]))
mod = Int8Linear.from_linear(lin)
deq = mod.weight.float() * mod.scale[:, None]
for row in range(out_f):
if row == 2:
assert torch.equal(mod.weight[row], torch.zeros(in_f, dtype=torch.int8))
assert float(mod.scale[row]) == 1.0
assert torch.equal(deq[row], torch.zeros(in_f))
else:
assert torch.equal(deq[row], weight[row]), (deq[row] - weight[row]).abs().max().item()
assert mod.bias is not None and torch.equal(mod.bias, lin.bias)
assert mod.bias.dtype == lin.bias.dtype
assert torch.isfinite(mod.scale).all()
# Random weights: per-element error stays within half a bin (+ float slack).
lin_r = nn.Linear(13, 9, bias=False)
mod_r = Int8Linear.from_linear(lin_r)
w = lin_r.weight.detach().float()
deq_r = (mod_r.weight.double() * mod_r.scale.double()[:, None]).float()
err = (w.double() - deq_r.double()).abs()
half = mod_r.scale.double()[:, None] * 0.5
slip = (err - half).max().item()
assert slip <= 1e-4, slip
assert torch.isfinite(mod_r.scale).all()
# Entirely zero weight: finite forward, zero codes, scale 1.
lin_z = nn.Linear(4, 3, bias=True)
with torch.no_grad():
lin_z.weight.zero_()
mod_z = Int8Linear.from_linear(lin_z)
assert torch.equal(mod_z.weight, torch.zeros_like(mod_z.weight))
assert torch.equal(mod_z.scale, torch.ones(3))
y = mod_z(torch.randn(8, 4))
assert torch.isfinite(y).all()
assert torch.allclose(y, mod_z.bias.expand_as(y))
# Zero row contributes only its bias.
x = torch.randn(6, in_f)
y_mix = mod(x)
assert torch.isfinite(y_mix).all()
assert torch.allclose(y_mix[:, 2], mod.bias[2].expand(6))
# bf16 source linear: codes int8, scale fp32, bias stays bf16.
lin_b = nn.Linear(8, 4, bias=True).to(dtype=torch.bfloat16)
mod_b = Int8Linear.from_linear(lin_b)
assert mod_b.weight.dtype == torch.int8
assert mod_b.scale.dtype == torch.float32
assert mod_b.bias is not None and mod_b.bias.dtype == torch.bfloat16
w_b = lin_b.weight.detach().float()
deq_b = mod_b.weight.float() * mod_b.scale[:, None]
err_b = (w_b.double() - deq_b.double()).abs()
half_b = mod_b.scale.double()[:, None] * 0.5
assert (err_b - half_b).max().item() <= 1e-2, (err_b - half_b).max().item()
def _assert_quant_dtypes(mod: Int8Linear, scale: torch.Tensor, weight: torch.Tensor, bias_dtype: torch.dtype) -> None:
assert mod.weight.dtype == torch.int8
assert mod.scale.dtype == torch.float32
assert torch.equal(mod.weight, weight)
assert torch.equal(mod.scale, scale)
assert mod.bias is not None and mod.bias.dtype == bias_dtype
def test_dtype_cast_keeps_scale_fp32() -> None:
torch.manual_seed(1)
lin = nn.Linear(5, 3, bias=True)
fresh = Int8Linear.from_linear(lin)
scale = fresh.scale.detach().clone()
weight = fresh.weight.detach().clone()
bias = fresh.bias.detach().clone()
assert scale.dtype == torch.float32 and weight.dtype == torch.int8 and bias.dtype == torch.float32
# Each cast starts from fp32 so "bias follows the cast" is the single cast of the source bias.
mod = Int8Linear.from_linear(lin)
mod.bfloat16()
_assert_quant_dtypes(mod, scale, weight, torch.bfloat16)
assert torch.equal(mod.bias, bias.to(dtype=torch.bfloat16))
mod = Int8Linear.from_linear(lin)
mod.half()
_assert_quant_dtypes(mod, scale, weight, torch.float16)
assert torch.equal(mod.bias, bias.to(dtype=torch.float16))
mod = Int8Linear.from_linear(lin)
mod.to(torch.bfloat16)
_assert_quant_dtypes(mod, scale, weight, torch.bfloat16)
assert torch.equal(mod.bias, bias.to(dtype=torch.bfloat16))
mod = Int8Linear.from_linear(lin)
mod.to(dtype=torch.float16)
_assert_quant_dtypes(mod, scale, weight, torch.float16)
assert torch.equal(mod.bias, bias.to(dtype=torch.float16))
# A second cast applies to the bias's current dtype, not the original fp32 value.
mod = Int8Linear.from_linear(lin)
mod.to(torch.bfloat16)
mod.to(dtype=torch.float16)
_assert_quant_dtypes(mod, scale, weight, torch.float16)
assert torch.equal(mod.bias, bias.to(dtype=torch.bfloat16).to(dtype=torch.float16))
# What the caller actually does: parent.to(device=..., dtype=bf16).
parent = nn.Sequential(Int8Linear.from_linear(lin))
parent.to(device="cpu", dtype=torch.bfloat16)
_assert_quant_dtypes(parent[0], scale, weight, torch.bfloat16)
assert torch.equal(parent[0].bias, bias.to(dtype=torch.bfloat16))
shell = Int8Linear.shell(4, 3, bias=True, bias_dtype=torch.bfloat16, device="meta")
assert shell.weight.dtype == torch.int8 and shell.weight.device.type == "meta"
assert shell.scale.dtype == torch.float32 and shell.scale.device.type == "meta"
assert shell.bias is not None
assert shell.bias.dtype == torch.bfloat16 and shell.bias.device.type == "meta"
shell_nb = Int8Linear.shell(4, 3, bias=False, bias_dtype=torch.float32, device="meta")
assert shell_nb.bias is None
def test_forward_matches_reference() -> None:
torch.manual_seed(2)
for bias in (True, False):
lin = nn.Linear(6, 4, bias=bias)
# Bias is passed through unchanged, so it has to already match x's dtype
# (the caller does model.to(dtype=...) before the prefill).
modules = [
(Int8Linear.from_linear(lin), torch.float32),
(Int8Linear.from_linear(lin).to(torch.bfloat16), torch.bfloat16),
(Int8Linear.from_linear(lin).to(dtype=torch.float16), torch.float16),
]
for mod, dtype in modules:
if mod.bias is not None:
assert mod.bias.dtype == dtype
x = torch.randn(3, 5, 6, dtype=dtype)
ref_w = (mod.weight.float() * mod.scale[:, None]).to(dtype=x.dtype)
y = mod(x)
y_ref = F.linear(x, ref_w, mod.bias)
assert torch.equal(y, y_ref), (bias, dtype)
def test_is_quantizable_rule() -> None:
false_cases = [
("model.model.layers.1.mlp.gate.weight", (256, 2048)),
("model.model.layers.1.mlp.image_gate.weight", (256, 2048)),
("model.model.layers.1.mlp.audio_gate.weight", (256, 2048)),
("model.model.layers.1.mlp.gate.expert_bias", (256,)),
("model.lm_head.weight", (151936, 2048)),
("model.model.word_embeddings.weight", (151936, 2048)),
("vision.blocks.0.attn.qkv.weight", (3072, 1280)),
("model.model.layers.0.input_layernorm.weight", (2048,)),
("model.model.layers.0.post_attention_layernorm.weight", (2048,)),
("model.model.layers.0.attention.q_norm.weight", (128,)),
("model.model.layers.0.attention.k_norm.weight", (128,)),
("model.model.norm.weight", (2048,)),
("linear_proj.0.weight", (2048, 2048)),
("model.model.layers.0.attention.query_key_value.bias", (3072,)),
("model.model.layers.0.mlp.experts.0.gate_proj.bias", (512,)),
# Right leaf, wrong rank: not quantizable (the stream must reject it).
("model.model.layers.0.attention.query_key_value.weight", (3072,)),
("model.model.layers.0.mlp.gate_proj.weight", (1024, 2048, 1)),
]
true_cases = [
("model.model.layers.3.mlp.experts.3.gate_proj.weight", (512, 2048)),
("model.model.layers.3.mlp.shared_experts.down_proj.weight", (2048, 512)),
("model.model.layers.0.mlp.up_proj.weight", (512, 2048)),
("layers.0.mlp.up_proj.weight", (512, 2048)),
("model.model.layers.0.attention.query_key_value.weight", (3072, 2048)),
("model.model.layers.0.attention.dense.weight", (2048, 2048)),
("model.model.layers.0.mlp.gate_proj.weight", (512, 2048)),
("model.model.layers.0.mlp.down_proj.weight", (2048, 512)),
("model.model.layers.19.mlp.experts.255.up_proj.weight", (512, 2048)),
]
for name, shape in false_cases:
assert is_quantizable(name, shape) is False, name
for name, shape in true_cases:
assert is_quantizable(name, shape) is True, name
def _shard_groups(sd: dict[str, torch.Tensor]) -> set[str]:
"""One copy-tensor, or one weight+scale pair, is one unsplittable group."""
names = set(sd)
groups: set[str] = set()
for name in names:
if name.endswith(".scale") and name[: -len(".scale")] + ".weight" in names:
groups.add(name[: -len(".scale")])
elif name.endswith(".weight") and name[: -len(".weight")] + ".scale" in names:
groups.add(name[: -len(".weight")])
else:
groups.add(name)
return groups
def test_end_to_end_stream_and_load() -> None:
assert quantize_stream.MAX_SHARD_BYTES == 5 * 10**9
torch.manual_seed(3)
src_model = TinyMing().to(dtype=torch.bfloat16)
# Non-persistent rotary buffer is not part of the checkpoint.
assert "model.model.layers.0.attention.inv_freq" not in src_model.state_dict()
with tempfile.TemporaryDirectory(prefix="ming-int8-") as tmp:
root = Path(tmp)
src = root / "src"
dst = root / "dst"
_save_bf16_checkpoint(src_model, src)
limit = 2048
old = quantize_stream.MAX_SHARD_BYTES
quantize_stream.MAX_SHARD_BYTES = limit
try:
rc = quantize_stream.main([str(src), str(dst)])
finally:
quantize_stream.MAX_SHARD_BYTES = old
assert rc == 0, rc
assert quantize_stream.MAX_SHARD_BYTES == 5 * 10**9
# Sidecars copied verbatim; original index replaced.
assert (dst / "config.json").read_bytes() == (src / "config.json").read_bytes()
assert (dst / "extra" / "chat_template.jinja").read_bytes() == (
src / "extra" / "chat_template.jinja"
).read_bytes()
assert not (dst / "bf16-00001.safetensors").exists()
manifest = json.loads((dst / "int8_manifest.json").read_text(encoding="utf-8"))
assert manifest["format"] == "ming-int8-wo-v1"
assert manifest["scheme"] == (
"weight-only int8, per-output-channel symmetric, fp32 scales"
)
assert manifest["quantized_modules"] == sorted(EXPECTED_QUANT_MODULES)
for banned in MUST_NOT_QUANTIZE:
assert banned not in manifest["quantized_modules"], banned
index = json.loads((dst / "model.safetensors.index.json").read_text(encoding="utf-8"))
assert index["metadata"]["total_size"] == manifest["total_size"]
measured = manifest["measured"]
assert measured["tensors_quantized"] == len(EXPECTED_QUANT_MODULES)
assert measured["bytes_in"] == manifest["source_total_size"]
assert measured["bytes_out"] == manifest["total_size"]
assert measured["bytes_out"] < measured["bytes_in"]
n_out_keys = measured["tensors_copied"] + 2 * measured["tensors_quantized"]
assert len(index["weight_map"]) == n_out_keys
src_sd = _load_all(src)
dst_sd = _load_all(dst)
assert manifest["source_total_size"] == sum(
t.numel() * t.element_size() for t in src_sd.values()
)
assert manifest["total_size"] == sum(t.numel() * t.element_size() for t in dst_sd.values())
shard_names = sorted({*index["weight_map"].values()})
assert len(shard_names) >= 2, shard_names
for shard in shard_names:
shard_sd = load_file(str(dst / shard))
total = sum(t.numel() * t.element_size() for t in shard_sd.values())
if total > limit:
assert len(_shard_groups(shard_sd)) == 1, (shard, total, list(shard_sd))
errors = []
for name, src_t in src_sd.items():
if is_quantizable(name, tuple(src_t.shape)):
q = dst_sd[name]
scale_key = name[: -len("weight")] + "scale"
scale = dst_sd[scale_key]
assert q.dtype == torch.int8, name
assert scale.dtype == torch.float32, scale_key
q_ref, scale_ref = quantize_weight(src_t)
assert torch.equal(q, q_ref), name
assert torch.equal(scale, scale_ref), scale_key
errors.append((name, quantize_stream._relative_frobenius(src_t, q, scale)))
else:
assert name in dst_sd, name
assert dst_sd[name].dtype == src_t.dtype, (name, dst_sd[name].dtype, src_t.dtype)
assert torch.equal(dst_sd[name], src_t), name
# Router weights stayed BF16 and byte-identical (the gate vs gate_proj trap).
router = "model.model.layers.1.mlp.gate.weight"
assert dst_sd[router].dtype == torch.bfloat16
assert torch.equal(dst_sd[router], src_sd[router])
for suffix in ("image_gate.weight", "audio_gate.weight", "gate.expert_bias"):
key = f"model.model.layers.1.mlp.{suffix}"
assert torch.equal(dst_sd[key], src_sd[key]), key
vals = [e for _, e in errors]
assert measured["max_relative_error"] == max(vals)
assert measured["mean_relative_error"] == sum(vals) / len(vals)
assert measured["worst_tensor"] in dict(errors)
assert measured["max_relative_error"] == dict(errors)[measured["worst_tensor"]]
assert 0.0 <= measured["mean_relative_error"] <= measured["p99_relative_error"]
assert measured["p99_relative_error"] <= measured["max_relative_error"]
assert measured["max_relative_error"] < 0.05, measured
# Eager quant of the same BF16 bytes.
eager = TinyMing().to(dtype=torch.bfloat16)
incompatible = eager.load_state_dict(src_sd, strict=True)
assert not incompatible.missing_keys and not incompatible.unexpected_keys
_apply_int8_(eager)
loaded = _move_parameters_to_meta(TinyMing())
for layer in loaded.model.model.layers:
assert layer.attention.inv_freq.device.type == "cpu"
assert layer.attention.query_key_value.weight.device.type == "meta"
report = load_int8_mllm_(loaded, dst, "cpu")
assert report["modules_swapped"] == len(EXPECTED_QUANT_MODULES)
assert report["tensors_loaded"] == len(dst_sd)
assert report["bytes_loaded"] == manifest["total_size"]
_assert_no_meta(loaded)
for layer in loaded.model.model.layers:
assert layer.attention.inv_freq.device.type == "cpu"
assert layer.attention.inv_freq.dtype == torch.float32
for name in EXPECTED_QUANT_MODULES:
mod = loaded.get_submodule(name)
assert isinstance(mod, Int8Linear), name
assert mod.weight.dtype == torch.int8
assert mod.scale.dtype == torch.float32
eager.eval()
loaded.eval()
ids = torch.randint(0, VOCAB, (2, 6))
with torch.no_grad():
y_eager = eager(ids)
y_loaded = loaded(ids)
assert y_eager.dtype == y_loaded.dtype
assert torch.equal(y_eager, y_loaded), (y_eager - y_loaded).abs().max().item()
# A second run into a non-empty safetensors dir must fail loudly.
print(" re-running into a non-empty dst (expect error on stderr)", flush=True)
rc_again = quantize_stream.main([str(src), str(dst)])
assert rc_again == 1
def test_unknown_key_fails_loudly() -> None:
torch.manual_seed(4)
model = TinyMing().to(dtype=torch.bfloat16)
with tempfile.TemporaryDirectory(prefix="ming-int8-bad-") as tmp:
root = Path(tmp)
src = root / "src"
dst = root / "dst"
_save_bf16_checkpoint(model, src)
rc = quantize_stream.main([str(src), str(dst)])
assert rc == 0, rc
shard = next(dst.glob("*.safetensors"))
sd = load_file(str(shard))
sd["not.a.real.key"] = torch.zeros(4, dtype=torch.float32)
save_file(sd, str(shard))
loaded = _move_parameters_to_meta(TinyMing())
try:
load_int8_mllm_(loaded, dst, "cpu")
except RuntimeError as exc:
text = str(exc)
assert "unexpected" in text.lower(), text
assert "not.a.real.key" in text, text
print(f" caught RuntimeError: {text.splitlines()[0]}")
else:
raise AssertionError("load_int8_mllm_ returned instead of failing on an unknown key")
def main() -> int:
import safetensors
print(f"torch={torch.__version__} safetensors={safetensors.__version__}", flush=True)
tests = [
test_from_linear_roundtrip,
test_dtype_cast_keeps_scale_fp32,
test_forward_matches_reference,
test_is_quantizable_rule,
test_end_to_end_stream_and_load,
test_unknown_key_fails_loudly,
]
failed = 0
for fn in tests:
try:
fn()
except Exception:
failed += 1
print(f"FAIL {fn.__name__}", flush=True)
traceback.print_exc()
else:
print(f"PASS {fn.__name__}", flush=True)
print(f"{len(tests) - failed} passed, {failed} failed", flush=True)
return 1 if failed else 0
if __name__ == "__main__":
sys.exit(main())