__merge-mtpKEYS-with-finetune.py
1.7 KB · 48 lines · python Raw
1 import os
2 import json
3 from safetensors.torch import load_file, save_file
4
5 def fix_sharded_mtp(base_path, merged_path):
6 # Load indices
7 with open(os.path.join(base_path, "model.safetensors.index.json"), "r") as f:
8 base_idx = json.load(f)
9 with open(os.path.join(merged_path, "model.safetensors.index.json"), "r") as f:
10 merged_idx = json.load(f)
11
12 # Find missing MTP keys
13 mtp_keys = {k: v for k, v in base_idx["weight_map"].items() if "mtp" in k.lower()}
14 missing_keys = [k for k in mtp_keys if k not in merged_idx["weight_map"]]
15
16 if not missing_keys:
17 return print("No missing MTP keys found.")
18
19 # Group by source shard to minimize loads
20 shard_to_keys = {}
21 for k in missing_keys:
22 shard = mtp_keys[k]
23 shard_to_keys.setdefault(shard, []).append(k)
24
25 # Collect tensors
26 restored_tensors = {}
27 for shard, keys in shard_to_keys.items():
28 weights = load_file(os.path.join(base_path, shard))
29 for k in keys:
30 restored_tensors[k] = weights[k]
31
32 # Save to a new dedicated MTP shard
33 new_shard_name = "model-mtp-restored.safetensors"
34 save_file(restored_tensors, os.path.join(merged_path, new_shard_name))
35
36 # Update merged index
37 for k in missing_keys:
38 merged_idx["weight_map"][k] = new_shard_name
39
40 with open(os.path.join(merged_path, "model.safetensors.index.json"), "w") as f:
41 json.dump(merged_idx, f, indent=2)
42
43 print(f"Fixed! Added {len(missing_keys)} keys to {new_shard_name}")
44
45 #
46 # Usage: fix_sharded_mtp("path/to/base", "path/to/merged")
47
48 fix_sharded_mtp("MODEL_W_MTP_KEYS", "MODEL_TO_ADD_MTP_KEYS")