materialize_ternary_1b.py
| 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 | |