kgrabko commited on
Commit
e8ec6fa
·
verified ·
1 Parent(s): 3496c4e

Create merge_sft_sandwich.py

Browse files
Files changed (1) hide show
  1. merge_sft_sandwich.py +57 -0
merge_sft_sandwich.py ADDED
@@ -0,0 +1,57 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ #%%writefile merge_sft_v30.py
2
+ import torch, glob, os
3
+ from safetensors.torch import load_file, save_file
4
+ from tqdm import tqdm
5
+
6
+ # Настройки
7
+ BASE_MODEL_DIR = "/content/JiRack_BitNet_405B_Packed"
8
+ SFT_CHECKPOINT = "/content/drive/MyDrive/405b_sft_results/sft_405b_v30_step_250.safetensors"
9
+ OUTPUT_DIR = "/content/JiRack_405B_SFT_Merged"
10
+ os.makedirs(OUTPUT_DIR, exist_ok=True)
11
+
12
+ # Маппинг слоев (тот же, что был в v30)
13
+ layer_map = {0:0, 1:1, 2:124, 3:125}
14
+
15
+ def merge():
16
+ print(f"💉 Загружаем SFT веса из {SFT_CHECKPOINT}...")
17
+ sft_sd = load_file(SFT_CHECKPOINT)
18
+
19
+ shard_files = sorted(glob.glob(f"{BASE_MODEL_DIR}/**/*.safetensors", recursive=True))
20
+
21
+ for shard in tqdm(shard_files, desc="Merging shards"):
22
+ sd = load_file(shard)
23
+ new_sd = {}
24
+ changed = False
25
+
26
+ for name, param in sd.items():
27
+ # Проверяем, есть ли этот параметр в наших обученных весах
28
+ if name in sft_sd:
29
+ new_sd[name] = sft_sd[name].to(param.dtype)
30
+ changed = True
31
+ continue
32
+
33
+ # Проверяем маппинг слоев (для слоев 0, 1, 124, 125)
34
+ # В SFT весах они лежат как layers.0, layers.1, layers.2, layers.3
35
+ if "layers." in name:
36
+ parts = name.split('.')
37
+ old_idx = int(parts[1])
38
+ # Находим, какой индекс в SFT соответствует этому слою в базе
39
+ sft_idx = next((k for k, v in layer_map.items() if v == old_idx), None)
40
+
41
+ if sft_idx is not None:
42
+ sft_name = name.replace(f"layers.{old_idx}", f"layers.{sft_idx}")
43
+ if sft_name in sft_sd:
44
+ new_sd[name] = sft_sd[sft_name].to(param.dtype)
45
+ changed = True
46
+ continue
47
+
48
+ new_sd[name] = param
49
+
50
+ # Сохраняем обновленный шард
51
+ output_shard = os.path.join(OUTPUT_DIR, os.path.basename(shard))
52
+ save_file(new_sd, output_shard)
53
+
54
+ print(f"✅ Мердж завершен! Модель сохранена в {OUTPUT_DIR}")
55
+
56
+ if __name__ == "__main__":
57
+ merge()