"""Stream a Ming MLLM directory to weight-only INT8 shards. Never builds the model: it buffers at most one output shard (<= 5 GB) of tensors at a time. Measured on the real 34.0 GB checkpoint (AMD Strix Halo, 2026-09-23): 266 s wall, peak RSS 17.8 GiB. CLI: ``python quantize_stream.py SRC_MLLM_DIR DST_DIR [--exclude MODULE_REGEX]`` (matching modules stay BF16). """ from __future__ import annotations import json import math import re import os import shutil import sys from dataclasses import dataclass from pathlib import Path import torch from safetensors import safe_open from safetensors.torch import save_file try: # imported as the `quant` package from .int8_linear import QUANT_RULE, quant_rule_leaf, quantize_weight except ImportError: # run as a script: python quant/quantize_stream.py SRC DST from int8_linear import QUANT_RULE, quant_rule_leaf, quantize_weight # Decimal GB, same unit Hugging Face uses for max_shard_size="5GB". MAX_SHARD_BYTES = 5 * 10**9 _DTYPE_BYTES = { "BOOL": 1, "U8": 1, "I8": 1, "F8_E4M3": 1, "F8_E5M2": 1, "F8_E8M0": 1, "U16": 2, "I16": 2, "F16": 2, "BF16": 2, "U32": 4, "I32": 4, "F32": 4, "U64": 8, "I64": 8, "F64": 8, } INDEX_NAME = "model.safetensors.index.json" MANIFEST_NAME = "int8_manifest.json" class QuantizeError(Exception): """User-facing checkpoint error. main() prints it and returns 1.""" def _die(msg: str) -> None: raise QuantizeError(msg) def _normalize_dtype(dtype_name) -> str: text = str(dtype_name).upper() if "." in text: text = text.rsplit(".", 1)[-1] aliases = { "BFLOAT16": "BF16", "FLOAT16": "F16", "FLOAT32": "F32", "FLOAT64": "F64", "FLOAT8_E4M3FN": "F8_E4M3", "FLOAT8_E5M2": "F8_E5M2", "INT8": "I8", "INT16": "I16", "INT32": "I32", "INT64": "I64", "UINT8": "U8", } return aliases.get(text, text) def _dtype_nbytes(dtype_name: str) -> int: try: return _DTYPE_BYTES[dtype_name] except KeyError: _die(f"unsupported safetensors dtype {dtype_name!r}") raise # unreachable; satisfies type checkers def _numel(shape: tuple[int, ...]) -> int: n = 1 for d in shape: n *= int(d) return n def _load_index(path: Path) -> dict: if not path.is_file(): _die(f"missing index: {path}") def _pairs(pairs): keys = [k for k, _ in pairs] dupes = sorted({k for k in keys if keys.count(k) > 1}) if dupes: _die(f"duplicate key(s) in {path}: {dupes}") return dict(pairs) try: raw = path.read_text(encoding="utf-8") index = json.loads(raw, object_pairs_hook=_pairs) except QuantizeError: raise except (OSError, json.JSONDecodeError) as exc: _die(f"cannot read index {path}: {exc}") if not isinstance(index, dict) or not isinstance(index.get("weight_map"), dict): _die(f"index {path} has no weight_map object") if not index["weight_map"]: _die(f"index {path} weight_map is empty") return index def _check_dst_clean(dst: Path) -> None: if not dst.exists(): return if not dst.is_dir(): _die(f"destination is not a directory: {dst}") found = sorted(p.relative_to(dst).as_posix() for p in dst.rglob("*.safetensors")) if found: _die(f"destination already contains safetensors: {found}") def _reject_nested(src: Path, dst: Path) -> None: src_r = src.resolve() dst_r = dst.resolve() if src_r == dst_r or src_r in dst_r.parents or dst_r in src_r.parents: _die(f"SRC and DST must be distinct and not nested: {src} vs {dst}") def _shard_path(src: Path, shard_name: str) -> Path: rel = Path(shard_name) if rel.is_absolute() or ".." in rel.parts: _die(f"unsafe shard path in index: {shard_name}") path = src / rel if not path.is_file(): _die(f"index lists missing shard: {shard_name}") return path @dataclass class Item: src_shard: str name: str kind: str # "copy" or "quant" shape: tuple[int, ...] src_dtype: str src_bytes: int out_bytes: int group: int = -1 def _scale_name(weight_name: str) -> str: return weight_name[: -len("weight")] + "scale" def _plan(src: Path, index: dict, exclude: str | None = None) -> list[Item]: """Metadata-only pass. Reads shapes and dtypes, not tensor bodies.""" weight_map: dict[str, str] = index["weight_map"] shard_order: list[str] = [] seen_shards: set[str] = set() for shard in weight_map.values(): if shard not in seen_shards: seen_shards.add(shard) shard_order.append(shard) index_names_by_shard: dict[str, set[str]] = {s: set() for s in shard_order} for name, shard in weight_map.items(): if shard not in index_names_by_shard: _die(f"weight_map value {shard!r} for {name} was not collected") index_names_by_shard[shard].add(name) items: list[Item] = [] seen_names: dict[str, str] = {} for shard in shard_order: path = _shard_path(src, shard) with safe_open(str(path), framework="pt", device="cpu") as handle: file_names = list(handle.keys()) file_set = set(file_names) if len(file_set) != len(file_names): _die(f"shard {shard} header lists a tensor name twice") missing = sorted(index_names_by_shard[shard] - file_set) extra = sorted(file_set - index_names_by_shard[shard]) if missing: _die(f"index lists tensors missing from {shard}: {missing}") if extra: _die(f"{shard} contains tensors absent from the index: {extra}") for name in file_names: if name in seen_names: _die( f"tensor name appears twice: {name} " f"({seen_names[name]} and {shard})" ) seen_names[name] = shard sl = handle.get_slice(name) if not hasattr(sl, "get_dtype") or not hasattr(sl, "get_shape"): _die( "safetensors safe_open slice is missing get_shape/get_dtype; " "cannot plan shards without loading tensor bodies" ) shape = tuple(int(d) for d in sl.get_shape()) dtype_name = _normalize_dtype(sl.get_dtype()) src_bytes = _numel(shape) * _dtype_nbytes(dtype_name) leaf = quant_rule_leaf(name) if leaf is not None and exclude and re.search(exclude, name[: -len(".weight")]): leaf = None # kept BF16 by --exclude if leaf is not None and len(shape) != 2: _die( f"tensor {name} matches the quantize rule but is not 2-D " f"(shape={list(shape)}, dtype={dtype_name})" ) if leaf is not None: out_bytes = _numel(shape) * 1 + shape[0] * 4 # int8 weight + fp32 scale items.append( Item(shard, name, "quant", shape, dtype_name, src_bytes, out_bytes) ) else: items.append( Item(shard, name, "copy", shape, dtype_name, src_bytes, src_bytes) ) index_names = set(weight_map) planned = {it.name for it in items} if planned != index_names: _die( "index / shard mismatch after scan: " f"only_in_index={sorted(index_names - planned)[:8]} " f"only_in_shards={sorted(planned - index_names)[:8]}" ) produced = set(planned) for it in items: if it.kind != "quant": continue sname = _scale_name(it.name) if sname in produced: _die(f"scale name collides with an existing tensor: {sname}") produced.add(sname) return items def _assign_groups(items: list[Item], max_shard_bytes: int) -> list[list[Item]]: if max_shard_bytes <= 0: _die(f"max_shard_bytes must be positive, got {max_shard_bytes}") groups: list[list[Item]] = [] cur: list[Item] = [] cur_bytes = 0 for it in items: if cur and cur_bytes + it.out_bytes > max_shard_bytes: groups.append(cur) cur = [] cur_bytes = 0 if cur_bytes == 0 and it.out_bytes > max_shard_bytes: print( f"warning: {it.name} contributes {it.out_bytes} bytes, " f"over the {max_shard_bytes}-byte shard target; writing it alone", file=sys.stderr, flush=True, ) it.group = len(groups) cur.append(it) cur_bytes += it.out_bytes if cur: groups.append(cur) return groups def _relative_frobenius(weight: torch.Tensor, q: torch.Tensor, scale: torch.Tensor) -> float: w = weight.detach().to(dtype=torch.float64) deq = q.detach().to(dtype=torch.float64) * scale.detach().to(dtype=torch.float64)[:, None] denom = torch.linalg.matrix_norm(w, ord="fro") numer = torch.linalg.matrix_norm(w - deq, ord="fro") d = denom.item() n = numer.item() if d == 0.0: return 0.0 if n == 0.0 else math.inf return n / d def _percentile_linear(values: list[float], pct: float) -> float: """NumPy-style linear percentile. Empty → 0.""" if not values: return 0.0 ordered = sorted(values) if len(ordered) == 1: return ordered[0] rank = (len(ordered) - 1) * (pct / 100.0) lo = math.floor(rank) hi = math.ceil(rank) if lo == hi: return ordered[lo] w = rank - lo return ordered[lo] * (1.0 - w) + ordered[hi] * w def _copy_sidecars(src: Path, dst: Path) -> list[str]: copied: list[str] = [] for dirpath, _dirnames, filenames in os.walk(src): rel = Path(dirpath).relative_to(src) out_dir = dst / rel out_dir.mkdir(parents=True, exist_ok=True) for filename in filenames: if filename.endswith(".safetensors"): continue if filename == INDEX_NAME and rel == Path("."): continue src_file = Path(dirpath) / filename dst_file = out_dir / filename shutil.copy2(src_file, dst_file) copied.append((rel / filename).as_posix()) return copied def _write_shards( src: Path, dst: Path, items: list[Item], groups: list[list[Item]], ) -> tuple[dict[str, str], int, int, list[tuple[str, float]], list[Path]]: n_out = len(groups) weight_map: dict[str, str] = {} bytes_in = 0 bytes_out = 0 errors: list[tuple[str, float]] = [] written: list[Path] = [] n_src = len({it.src_shard for it in items}) src_seen = 0 open_name: str | None = None handle = None buf: dict[str, torch.Tensor] = {} buf_q = 0 buf_c = 0 current_group = 0 def flush() -> None: nonlocal buf, buf_q, buf_c, current_group if not buf: return fname = f"model-{current_group + 1:05d}-of-{n_out:05d}.safetensors" path = dst / fname for key, tensor in buf.items(): if not tensor.is_contiguous(): buf[key] = tensor.contiguous() save_file(buf, str(path)) shard_bytes = 0 for key, tensor in buf.items(): weight_map[key] = fname shard_bytes += tensor.numel() * tensor.element_size() written.append(path) print( f"wrote {fname}: tensors={len(buf)} quantized={buf_q} copied={buf_c} " f"bytes={shard_bytes}", flush=True, ) buf = {} buf_q = 0 buf_c = 0 current_group += 1 try: for it in items: if it.src_shard != open_name: if handle is not None: handle.__exit__(None, None, None) handle = None path = _shard_path(src, it.src_shard) handle = safe_open(str(path), framework="pt", device="cpu") handle.__enter__() open_name = it.src_shard src_seen += 1 n_here = sum(1 for x in items if x.src_shard == it.src_shard) print( f"reading source shard {src_seen}/{n_src} {it.src_shard} ({n_here} tensors)", flush=True, ) assert handle is not None tensor = handle.get_tensor(it.name) got = tensor.numel() * tensor.element_size() if got != it.src_bytes: _die( f"{it.name} byte size {got} != planned {it.src_bytes} " f"(dtype={tensor.dtype}, shape={tuple(tensor.shape)})" ) bytes_in += got if it.kind == "quant": if not tensor.is_floating_point(): _die( f"{it.name} matches the quantize rule but dtype is {tensor.dtype}, " "expected a floating dtype" ) if tuple(tensor.shape) != it.shape: _die(f"{it.name} shape changed between passes: {tuple(tensor.shape)} vs {it.shape}") q, scale = quantize_weight(tensor) err = _relative_frobenius(tensor, q, scale) if math.isnan(err) or math.isinf(err): _die(f"non-finite relative error for {it.name}: {err}") errors.append((it.name, err)) del tensor sname = _scale_name(it.name) buf[it.name] = q buf[sname] = scale produced = q.numel() * q.element_size() + scale.numel() * scale.element_size() if produced != it.out_bytes: _die(f"{it.name} output bytes {produced} != planned {it.out_bytes}") buf_q += 1 else: if not tensor.is_contiguous(): tensor = tensor.contiguous() buf[it.name] = tensor buf_c += 1 bytes_out += it.out_bytes # Flush when this item closes its planned output shard. group_items = groups[it.group] if it is group_items[-1]: flush() finally: if handle is not None: handle.__exit__(None, None, None) if buf: _die("internal error: output buffer not flushed") if current_group != n_out: _die(f"internal error: wrote {current_group} shards, planned {n_out}") return weight_map, bytes_in, bytes_out, errors, written def _summary( errors: list[tuple[str, float]], n_quant: int, n_copy: int, bytes_in: int, bytes_out: int, ) -> dict: vals = [e for _, e in errors] if errors: worst_name, worst_err = min( errors, key=lambda pair: (-pair[1], pair[0]), ) else: worst_name, worst_err = None, 0.0 mean = (sum(vals) / len(vals)) if vals else 0.0 return { "tensors_quantized": n_quant, "tensors_copied": n_copy, "bytes_in": bytes_in, "bytes_out": bytes_out, "mean_relative_error": mean, "p99_relative_error": _percentile_linear(vals, 99.0), "max_relative_error": worst_err if vals else 0.0, "worst_tensor": worst_name, } def run(src: Path, dst: Path, max_shard_bytes: int = MAX_SHARD_BYTES, exclude: str | None = None) -> dict: src = src.resolve() dst = dst.resolve() if not src.is_dir(): _die(f"SRC is not a directory: {src}") _reject_nested(src, dst) _check_dst_clean(dst) index = _load_index(src / INDEX_NAME) items = _plan(src, index, exclude) groups = _assign_groups(items, max_shard_bytes) dst.mkdir(parents=True, exist_ok=True) written: list[Path] = [] try: weight_map, bytes_in, bytes_out, errors, written = _write_shards(src, dst, items, groups) copied = _copy_sidecars(src, dst) n_quant = sum(1 for it in items if it.kind == "quant") n_copy = sum(1 for it in items if it.kind == "copy") measured = _summary(errors, n_quant, n_copy, bytes_in, bytes_out) if measured["bytes_in"] != bytes_in or measured["bytes_out"] != bytes_out: _die("internal error: summary byte counters diverged") # Recompute the on-disk total from the tensors we recorded. weight_map # values are what we just saved; bytes_out is that sum. out_index = {"metadata": {"total_size": bytes_out}, "weight_map": weight_map} (dst / INDEX_NAME).write_text( json.dumps(out_index, indent=2) + "\n", encoding="utf-8" ) modules = sorted( it.name[: -len(".weight")] for it in items if it.kind == "quant" ) manifest = { "format": "ming-int8-wo-v1", "scheme": "weight-only int8, per-output-channel symmetric, fp32 scales", "rule": QUANT_RULE + (f" Additionally kept BF16: modules matching /{exclude}/." if exclude else ""), "exclude": exclude, "quantized_modules": modules, "source_total_size": bytes_in, "total_size": bytes_out, "measured": measured, } (dst / MANIFEST_NAME).write_text( json.dumps(manifest, indent=2, allow_nan=False) + "\n", encoding="utf-8" ) except Exception: for path in written: try: path.unlink() except OSError: pass raise print(f"copied {len(copied)} non-safetensors file(s)", flush=True) m = measured print( "summary: " f"quantized={m['tensors_quantized']} copied={m['tensors_copied']} " f"bytes_in={m['bytes_in']} bytes_out={m['bytes_out']} " f"mean_rel={m['mean_relative_error']:.8g} " f"p99_rel={m['p99_relative_error']:.8g} " f"max_rel={m['max_relative_error']:.8g} " f"worst={m['worst_tensor']}", flush=True, ) return manifest def main(argv: list[str] | None = None) -> int: args = list(sys.argv[1:] if argv is None else argv) exclude = None if "--exclude" in args: i = args.index("--exclude") if i + 1 >= len(args): print("--exclude needs a regex", file=sys.stderr) return 2 exclude = args[i + 1] re.compile(exclude) del args[i : i + 2] if len(args) != 2: print( "usage: python quantize_stream.py SRC_MLLM_DIR DST_DIR [--exclude MODULE_REGEX]", file=sys.stderr, ) return 2 try: run(Path(args[0]), Path(args[1]), max_shard_bytes=MAX_SHARD_BYTES, exclude=exclude) except QuantizeError as exc: print(f"error: {exc}", file=sys.stderr) return 1 return 0 if __name__ == "__main__": sys.exit(main())