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")