diff --git a/scripts/strip_vision_weights.py b/scripts/strip_vision_weights.py index ce87800..1990855 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,16 @@ 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 + # index: remove vision keys and recompute total_size by the same + # amount actually dropped from the shards above (untouched shards' + # bytes are unaffected, so this equals recomputing from scratch) idx["weight_map"] = {k: v for k, v in wm.items() if k not in vs} + if "total_size" in idx.get("metadata", {}): + idx["metadata"]["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..1b12da8 --- /dev/null +++ b/tests/test_strip_vision_weights.py @@ -0,0 +1,87 @@ +"""scripts/strip_vision_weights.py: index bookkeeping after stripping. + +Builds a minimal two-tensor safetensors shard (one vision, one text) by +hand so the test needs neither MLX nor a real checkpoint. +""" + +from __future__ import annotations + +import importlib.util +import json +import struct +import sys +from pathlib import Path + +ROOT = Path(__file__).resolve().parents[1] + + +def _load_strip_vision_weights(): + spec = importlib.util.spec_from_file_location( + "strip_vision_weights", ROOT / "scripts" / "strip_vision_weights.py") + mod = importlib.util.module_from_spec(spec) + sys.modules[spec.name] = mod + spec.loader.exec_module(mod) + return mod + + +def _write_shard(path: Path, hdr: dict, data: bytes) -> None: + hdr_bytes = json.dumps(hdr).encode() + pad = (8 - (len(hdr_bytes) % 8)) % 8 + hdr_bytes += b" " * pad + with open(path, "wb") as f: + f.write(struct.pack(" None: + # 4 bytes of vision tensor, 4 bytes of text tensor. + hdr = { + "__metadata__": {"format": "pt"}, + "vision_tower.a": {"dtype": "F32", "shape": [1], + "data_offsets": [0, 4]}, + "model.layers.0.w": {"dtype": "F32", "shape": [1], + "data_offsets": [4, 8]}, + } + _write_shard(model_dir / "model.safetensors", hdr, b"\x00" * 8) + idx = { + "metadata": {"total_size": 8}, + "weight_map": { + "vision_tower.a": "model.safetensors", + "model.layers.0.w": "model.safetensors", + }, + } + (model_dir / "model.safetensors.index.json").write_text(json.dumps(idx)) + + +def test_total_size_shrinks_by_dropped_bytes(tmp_path, monkeypatch): + _make_checkpoint(tmp_path) + mod = _load_strip_vision_weights() + monkeypatch.setattr( + sys, "argv", ["strip_vision_weights.py", str(tmp_path)]) + assert mod.main() == 0 + + idx = json.loads( + (tmp_path / "model.safetensors.index.json").read_text()) + assert "vision_tower.a" not in idx["weight_map"] + assert idx["metadata"]["total_size"] == 4 # 8 - 4 dropped vision bytes + + +def test_no_vision_weights_leaves_total_size_untouched(tmp_path, monkeypatch): + hdr = {"model.layers.0.w": {"dtype": "F32", "shape": [1], + "data_offsets": [0, 4]}} + _write_shard(tmp_path / "model.safetensors", hdr, b"\x00" * 4) + idx = { + "metadata": {"total_size": 4}, + "weight_map": {"model.layers.0.w": "model.safetensors"}, + } + (tmp_path / "model.safetensors.index.json").write_text(json.dumps(idx)) + + mod = _load_strip_vision_weights() + monkeypatch.setattr( + sys, "argv", ["strip_vision_weights.py", str(tmp_path)]) + assert mod.main() == 0 + + idx_after = json.loads( + (tmp_path / "model.safetensors.index.json").read_text()) + assert idx_after == idx # untouched: nothing to strip