Instructions to use kingjones777/Ming-Image-0.1-Design-ROCm-INT8 with libraries, inference providers, notebooks, and local apps. Follow these links to get started.
- Libraries
- Diffusers
How to use kingjones777/Ming-Image-0.1-Design-ROCm-INT8 with Diffusers:
pip install -U diffusers transformers accelerate
import torch from diffusers import DiffusionPipeline # switch to "mps" for apple devices pipe = DiffusionPipeline.from_pretrained("kingjones777/Ming-Image-0.1-Design-ROCm-INT8", dtype=torch.bfloat16, device_map="cuda") prompt = "Astronaut in a jungle, cold color palette, muted colors, detailed, 8k" image = pipe(prompt).images[0] - Notebooks
- Google Colab
- Kaggle
- Local Apps Settings
- Draw Things
- DiffusionBee
Download code/quant/test_int8.py from kingjones777/Ming-Image-0.1-Design-ROCm-INT8: direct link, hf CLI and curl.
- Browser
- Download file 27.6 kB
-
https://huggingface.co/kingjones777/Ming-Image-0.1-Design-ROCm-INT8/resolve/main/code/quant/test_int8.py
- Command line
-
hf download hf://kingjones777/Ming-Image-0.1-Design-ROCm-INT8/code/quant/test_int8.py
-
curl -L -o test_int8.py https://huggingface.co/kingjones777/Ming-Image-0.1-Design-ROCm-INT8/resolve/main/code/quant/test_int8.py
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()) | |