#!/usr/bin/env python # ============================================================================= # materialize_ternary_1b.py # ============================================================================= # Мостик между обученным model.pt (JiRack-1b) и convert_hf_to_gguf.py. # # ПРОБЛЕМА: # BitLinear.forward() при lambda=1 вычисляет тернарные веса НА ЛЕТУ: # gamma = w.abs().mean().clamp(min=eps) # w_ternary = clamp(round(w / gamma), -1, 1) * gamma # Но в state_dict хранится СЫРОЙ обучаемый параметр w, а не w_ternary. # jirack_to_gguf_1p5b.py копирует именно сырой w -- то есть по факту # экспортирует плотную модель, даже если обучение дошло до lambda=1. # # ЧТО ДЕЛАЕТ СКРИПТ: # Загружает state_dict из model.pt, для каждого слоя, патченного как # BitLinear (см. PATCHED_LAYER_PATTERNS), применяет ту же формулу # gamma/round/clamp и заменяет вес на точное тернарное значение. # Всё остальное (embeddings, lm_head, нормы, bias) копируется как есть. # # Результат -- новый .pt, который подаётся в jirack_to_gguf_1p5b.py # вместо оригинального model.pt. После него convert_hf_to_gguf.py # --outtype bf16 даёт файл, для которого llama-quantize ... TQ2_0 # будет лоссless round-trip. # # ЧЕГО СКРИПТ НЕ ДЕЛАЕТ: # - не запускает обучение/warmup; # - не запускает convert_hf_to_gguf.py сам -- это следующий шаг руками; # - не чинит именование Q2_0 -> TQ2_0 в jirack_to_gguf_1p5b.py -- это # отдельная правка в том скрипте. # ============================================================================= import argparse import re import torch # Подстроить под реальные имена в твоём state_dict, если отличаются. # Проверить: python -c "import torch; print(list(torch.load('model.pt', map_location='cpu').keys()))" PATCHED_LAYER_PATTERNS = [ r"\.q_proj\.weight$", r"\.k_proj\.weight$", r"\.v_proj\.weight$", r"\.o_proj\.weight$", r"\.gate_proj\.weight$", r"\.up_proj\.weight$", r"\.down_proj\.weight$", r"\.ffn_w1\.weight$", r"\.ffn_w2\.weight$", r"\.ffn_w3\.weight$", r"\.out_proj\.weight$", ] EPS = 1e-5 # тот же eps, что в BitLinear -- подстроить, если у тебя другой def is_patched_layer(key: str) -> bool: return any(re.search(pat, key) for pat in PATCHED_LAYER_PATTERNS) def ternarize(w: torch.Tensor): """Точная копия формулы из BitLinear.forward() при lambda=1.""" w32 = w.float() gamma = w32.abs().mean().clamp(min=EPS) w_ternary = torch.clamp(torch.round(w32 / gamma), -1, 1) * gamma return w_ternary.to(w.dtype), gamma.item() def main(): ap = argparse.ArgumentParser(description=__doc__) ap.add_argument("input_pt", help="путь к обученному model.pt") ap.add_argument("output_pt", help="куда сохранить материализованную тернарную версию") ap.add_argument("--dry-run", action="store_true", help="только показать, что будет дискретизировано, ничего не сохранять") args = ap.parse_args() print(f"📥 Загрузка {args.input_pt} ...") sd = torch.load(args.input_pt, map_location="cpu") outer = None if isinstance(sd, dict) and "state_dict" in sd and not any(k.endswith(".weight") for k in sd.keys()): print(" обнаружена обёртка с ключом 'state_dict', разворачиваю") outer = sd sd = outer["state_dict"] total = 0 ternarized = 0 max_relative_shift = 0.0 for key, tensor in sd.items(): if not torch.is_tensor(tensor) or not tensor.dtype.is_floating_point: continue total += 1 if not is_patched_layer(key): continue w_new, gamma = ternarize(tensor) shift = (w_new.float() - tensor.float()).abs().max().item() scale = tensor.float().abs().max().item() rel = shift / scale if scale > 0 else 0.0 max_relative_shift = max(max_relative_shift, rel) print(f" {key:60s} gamma={gamma:.6f} max|Δ|={shift:.6f} (отн. {rel:.2%})") if not args.dry_run: sd[key] = w_new ternarized += 1 print(f"\n✅ Тензоров всего: {total}, дискретизировано: {ternarized}") print(f" Максимальный относительный сдвиг веса: {max_relative_shift:.2%}") if max_relative_shift > 0.15: print(" ⚠️ Сдвиг заметный (>15%) -- вероятно, warmup не дошёл до " "lambda=1 или дообучение было коротким. После дискретизации " "качество может заметно просесть. Стоит сверить лог обучения.") else: print(" Сдвиг небольшой -- веса уже были близки к тернарным, " "дискретизация должна пройти почти без потери качества.") if args.dry_run: print("\n(dry-run: файл не сохранён)") return if outer is not None: outer["state_dict"] = sd torch.save(outer, args.output_pt) else: torch.save(sd, args.output_pt) print(f"\n💾 Сохранено: {args.output_pt}") print(" Дальше: подать этот файл в jirack_to_gguf_1p5b.py вместо " "исходного model.pt, затем как обычно convert_hf_to_gguf.py " "--outtype bf16, и llama-quantize ... TQ2_0 (не Q2_0 -- в " "llama.cpp нет типа с таким именем).") if __name__ == "__main__": main()