File size: 1,450 Bytes
fdc6474 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 | #!/usr/bin/env python3
"""Physically remove tensors not referenced by the index from /data/glm52.
vLLM's local-dir loader iterates every tensor in every *.safetensors file,
ignoring the index, so de-indexed tensors must actually be deleted. Rewrites
any shard containing unreferenced keys; deletes shards left empty.
"""
import json
import os
import sys
from safetensors import safe_open
from safetensors.torch import save_file
DST = sys.argv[1] if len(sys.argv) > 1 else "/data/glm52"
idx = json.load(open(os.path.join(DST, "model.safetensors.index.json")))
wm = idx["weight_map"]
files = sorted(f for f in os.listdir(DST) if f.endswith(".safetensors"))
total_removed = 0
for f in files:
path = os.path.join(DST, f)
with safe_open(path, framework="pt") as sf:
keys = list(sf.keys())
keep = [k for k in keys if wm.get(k) == f]
if len(keep) == len(keys):
continue
removed = len(keys) - len(keep)
if not keep:
print(f"{f}: all {len(keys)} tensors stale -> deleting file")
tensors = None
else:
print(f"{f}: removing {removed}/{len(keys)} stale tensors")
tensors = {k: sf.get_tensor(k) for k in keep}
total_removed += removed
if tensors is None:
os.remove(path)
else:
tmp = path + ".tmp"
save_file(tensors, tmp)
os.replace(tmp, path)
print(f"done: removed {total_removed} stale tensors")
|