materialize_ternary_1b.py
6.2 KB · 136 lines · python Raw
1 #!/usr/bin/env python
2 # =============================================================================
3 # materialize_ternary_1b.py
4 # =============================================================================
5 # Мостик между обученным model.pt (JiRack-1b) и convert_hf_to_gguf.py.
6 #
7 # ПРОБЛЕМА:
8 # BitLinear.forward() при lambda=1 вычисляет тернарные веса НА ЛЕТУ:
9 # gamma = w.abs().mean().clamp(min=eps)
10 # w_ternary = clamp(round(w / gamma), -1, 1) * gamma
11 # Но в state_dict хранится СЫРОЙ обучаемый параметр w, а не w_ternary.
12 # jirack_to_gguf_1p5b.py копирует именно сырой w -- то есть по факту
13 # экспортирует плотную модель, даже если обучение дошло до lambda=1.
14 #
15 # ЧТО ДЕЛАЕТ СКРИПТ:
16 # Загружает state_dict из model.pt, для каждого слоя, патченного как
17 # BitLinear (см. PATCHED_LAYER_PATTERNS), применяет ту же формулу
18 # gamma/round/clamp и заменяет вес на точное тернарное значение.
19 # Всё остальное (embeddings, lm_head, нормы, bias) копируется как есть.
20 #
21 # Результат -- новый .pt, который подаётся в jirack_to_gguf_1p5b.py
22 # вместо оригинального model.pt. После него convert_hf_to_gguf.py
23 # --outtype bf16 даёт файл, для которого llama-quantize ... TQ2_0
24 # будет лоссless round-trip.
25 #
26 # ЧЕГО СКРИПТ НЕ ДЕЛАЕТ:
27 # - не запускает обучение/warmup;
28 # - не запускает convert_hf_to_gguf.py сам -- это следующий шаг руками;
29 # - не чинит именование Q2_0 -> TQ2_0 в jirack_to_gguf_1p5b.py -- это
30 # отдельная правка в том скрипте.
31 # =============================================================================
32
33 import argparse
34 import re
35 import torch
36
37 # Подстроить под реальные имена в твоём state_dict, если отличаются.
38 # Проверить: python -c "import torch; print(list(torch.load('model.pt', map_location='cpu').keys()))"
39 PATCHED_LAYER_PATTERNS = [
40 r"\.q_proj\.weight$",
41 r"\.k_proj\.weight$",
42 r"\.v_proj\.weight$",
43 r"\.o_proj\.weight$",
44 r"\.gate_proj\.weight$",
45 r"\.up_proj\.weight$",
46 r"\.down_proj\.weight$",
47 r"\.ffn_w1\.weight$",
48 r"\.ffn_w2\.weight$",
49 r"\.ffn_w3\.weight$",
50 r"\.out_proj\.weight$",
51 ]
52
53 EPS = 1e-5 # тот же eps, что в BitLinear -- подстроить, если у тебя другой
54
55
56 def is_patched_layer(key: str) -> bool:
57 return any(re.search(pat, key) for pat in PATCHED_LAYER_PATTERNS)
58
59
60 def ternarize(w: torch.Tensor):
61 """Точная копия формулы из BitLinear.forward() при lambda=1."""
62 w32 = w.float()
63 gamma = w32.abs().mean().clamp(min=EPS)
64 w_ternary = torch.clamp(torch.round(w32 / gamma), -1, 1) * gamma
65 return w_ternary.to(w.dtype), gamma.item()
66
67
68 def main():
69 ap = argparse.ArgumentParser(description=__doc__)
70 ap.add_argument("input_pt", help="путь к обученному model.pt")
71 ap.add_argument("output_pt", help="куда сохранить материализованную тернарную версию")
72 ap.add_argument("--dry-run", action="store_true",
73 help="только показать, что будет дискретизировано, ничего не сохранять")
74 args = ap.parse_args()
75
76 print(f"📥 Загрузка {args.input_pt} ...")
77 sd = torch.load(args.input_pt, map_location="cpu")
78
79 outer = None
80 if isinstance(sd, dict) and "state_dict" in sd and not any(k.endswith(".weight") for k in sd.keys()):
81 print(" обнаружена обёртка с ключом 'state_dict', разворачиваю")
82 outer = sd
83 sd = outer["state_dict"]
84
85 total = 0
86 ternarized = 0
87 max_relative_shift = 0.0
88
89 for key, tensor in sd.items():
90 if not torch.is_tensor(tensor) or not tensor.dtype.is_floating_point:
91 continue
92 total += 1
93 if not is_patched_layer(key):
94 continue
95
96 w_new, gamma = ternarize(tensor)
97 shift = (w_new.float() - tensor.float()).abs().max().item()
98 scale = tensor.float().abs().max().item()
99 rel = shift / scale if scale > 0 else 0.0
100 max_relative_shift = max(max_relative_shift, rel)
101
102 print(f" {key:60s} gamma={gamma:.6f} max|Δ|={shift:.6f} (отн. {rel:.2%})")
103
104 if not args.dry_run:
105 sd[key] = w_new
106 ternarized += 1
107
108 print(f"\n✅ Тензоров всего: {total}, дискретизировано: {ternarized}")
109 print(f" Максимальный относительный сдвиг веса: {max_relative_shift:.2%}")
110 if max_relative_shift > 0.15:
111 print(" ⚠️ Сдвиг заметный (>15%) -- вероятно, warmup не дошёл до "
112 "lambda=1 или дообучение было коротким. После дискретизации "
113 "качество может заметно просесть. Стоит сверить лог обучения.")
114 else:
115 print(" Сдвиг небольшой -- веса уже были близки к тернарным, "
116 "дискретизация должна пройти почти без потери качества.")
117
118 if args.dry_run:
119 print("\n(dry-run: файл не сохранён)")
120 return
121
122 if outer is not None:
123 outer["state_dict"] = sd
124 torch.save(outer, args.output_pt)
125 else:
126 torch.save(sd, args.output_pt)
127 print(f"\n💾 Сохранено: {args.output_pt}")
128 print(" Дальше: подать этот файл в jirack_to_gguf_1p5b.py вместо "
129 "исходного model.pt, затем как обычно convert_hf_to_gguf.py "
130 "--outtype bf16, и llama-quantize ... TQ2_0 (не Q2_0 -- в "
131 "llama.cpp нет типа с таким именем).")
132
133
134 if __name__ == "__main__":
135 main()
136