kingjones777's picture
Add files using upload-large-folder tool
18c1466 verified
Raw History Blame Contribute Delete
19.2 kB
"""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())