From 12635619286fd922b63985e66efe9cf2c1857a38 Mon Sep 17 00:00:00 2001 From: Cocoon-Break <54054995+kuishou68@users.noreply.github.com> Date: Sun, 27 Sep 2026 21:39:06 +0800 Subject: [PATCH] fix(strip_vision_weights): shrink metadata.total_size by dropped bytes After stripping vision-tower tensors, the script removed the dropped keys from model.safetensors.index.json but left metadata.total_size at its pre-strip value, so the index reported more bytes than the shards actually contain. Accumulate the dropped tensor bytes across shards and subtract them from metadata.total_size when the field is present. Add a stdlib-only regression test that builds a synthetic shard + index, runs the script, and checks the corrected total_size. Fixes #18 --- scripts/strip_vision_weights.py | 7 +++ tests/test_strip_vision_weights.py | 72 ++++++++++++++++++++++++++++++ 2 files changed, 79 insertions(+) create mode 100644 tests/test_strip_vision_weights.py diff --git a/scripts/strip_vision_weights.py b/scripts/strip_vision_weights.py index ce87800..ac1d630 100644 --- a/scripts/strip_vision_weights.py +++ b/scripts/strip_vision_weights.py @@ -41,6 +41,7 @@ def main() -> int: return 0 shards = sorted({wm[k] for k in vs}) + total_saved = 0 for shard_name in shards: path = d / shard_name hdr, data_start = read_shard(path) @@ -77,11 +78,17 @@ def main() -> int: f.write(out) saved = sum(hdr[k]["data_offsets"][1] - hdr[k]["data_offsets"][0] for k in drop) + total_saved += saved print(f"{shard_name}: dropped {len(drop)} tensors, " f"{saved/1e9:.2f} GB, new size {path.stat().st_size/1e9:.2f} GB") # index: remove vision keys idx["weight_map"] = {k: v for k, v in wm.items() if k not in vs} + md = idx.get("metadata") or {} + if "total_size" in md: + # metadata.total_size describes the pre-strip shards — shrink it by + # the tensor bytes actually dropped so it stays consistent + md["total_size"] -= total_saved if not args.no_backup: shutil.copy2(idx_path, idx_path.with_suffix(".json.bak_vision")) json.dump(idx, open(idx_path, "w"), indent=2) diff --git a/tests/test_strip_vision_weights.py b/tests/test_strip_vision_weights.py new file mode 100644 index 0000000..3aafb39 --- /dev/null +++ b/tests/test_strip_vision_weights.py @@ -0,0 +1,72 @@ +"""Regression test for scripts/strip_vision_weights.py (#18). + +Pure stdlib: builds a tiny synthetic safetensors shard + index, runs the +script as a subprocess, and checks that ``metadata.total_size`` shrinks +by the bytes actually dropped. Imports nothing from ``edge0`` so it runs +on any platform (including CI without MLX wheels). +""" +from __future__ import annotations + +import json +import struct +import subprocess +import sys +from pathlib import Path + +ROOT = Path(__file__).resolve().parents[1] +SCRIPT = ROOT / "scripts" / "strip_vision_weights.py" + + +def _write_shard(path: Path, tensors: dict) -> None: + header = {} + off = 0 + for name, blob in tensors.items(): + header[name] = {"dtype": "U8", "shape": [len(blob)], + "data_offsets": [off, off + len(blob)]} + off += len(blob) + header["__metadata__"] = {"format": "pt"} + hdr = json.dumps(header).encode() + hdr += b" " * ((8 - len(hdr) % 8) % 8) + with open(path, "wb") as f: + f.write(struct.pack(" dict: + with open(path, "rb") as f: + n = struct.unpack("