export_27b_to_hf_safetensors.py
5.2 KB · 121 lines · python Raw
1 """
2 export_27b_to_hf_safetensors.py
3 ================================
4 Обратный экспорт JiRack-чекпоинта (.pt) 27B в HF-совместимый
5 `model.safetensors` -- для последующего convert_hf_to_gguf.py.
6
7 MTP/NextN: JiRackDeltaNet_27b.py теперь содержит MTPBlock (top-level
8 `mtp.`), поэтому mtp.* тензоры попадают в экспорт. НО: старые JiRack
9 .pt-чекпоинты сохранены до появления MTPBlock и mtp-весов не содержат --
10 без --hf_mtp_dir голова уедет со случайной инициализацией.
11 """
12
13 import argparse
14 import os
15 import sys
16
17 import torch
18 from safetensors.torch import save_file
19
20 sys.path.insert(0, os.path.dirname(os.path.abspath(__file__)))
21 from JiRackDeltaNet_27b import JiRackQwen38ForCausalLM, BitLinear # noqa: E402
22
23
24 def main():
25 parser = argparse.ArgumentParser(description=__doc__)
26 parser.add_argument(
27 "--checkpoint", default="model.pt",
28 help="Путь к JiRack .pt чекпоинту (save_checkpoint формат)",
29 )
30 parser.add_argument(
31 "--out_dir", default="/mnt/nfs_share/Qwen3_8",
32 help="Куда положить model.safetensors (папка, не файл)",
33 )
34 parser.add_argument(
35 "--dtype", default="bfloat16", choices=["bfloat16", "float16", "float32"],
36 help="Тип тензоров на выходе",
37 )
38 parser.add_argument(
39 "--hf_mtp_dir", default=None,
40 help="Папка с ОРИГИНАЛЬНЫМ HF-чекпоинтом Qwen3.8-27B (шарды + "
41 "index.json), откуда подсадить mtp.* веса. Без этого флага "
42 "старый .pt даст СЛУЧАЙНУЮ MTP-голову!",
43 )
44 parser.add_argument(
45 "--include_gate_proj", action="store_true",
46 help="Учитывать ли in_proj_b/in_proj_a как патченные BitLinear "
47 "при загрузке. Игнорируется, если в чекпоинте уже сохранён "
48 "patched_linear_names.",
49 )
50 args = parser.parse_args()
51
52 dtype_map = {
53 "bfloat16": torch.bfloat16,
54 "float16": torch.float16,
55 "float32": torch.float32,
56 }
57 dtype = dtype_map[args.dtype]
58
59 print(f"📥 Загрузка JiRack-чекпоинта: {args.checkpoint}")
60 model = JiRackQwen38ForCausalLM.load_checkpoint(
61 args.checkpoint, dtype=dtype, include_gate_proj=args.include_gate_proj,
62 )
63 model.eval()
64
65 if args.hf_mtp_dir:
66 model.load_mtp_from_hf_dir(args.hf_mtp_dir)
67 else:
68 has_mtp = any(k.startswith("mtp.") for k in model.state_dict())
69 print(
70 "⚠️ --hf_mtp_dir не задан. Если .pt старый (без mtp.*), "
71 "MTP-голова уйдёт в safetensors со случайной инициализацией -- "
72 "GGUF соберётся, но спекулятивный декод будет мусорным."
73 if has_mtp else ""
74 )
75
76 lambdas = sorted({float(m.lambda_) for m in model.modules() if isinstance(m, BitLinear)})
77 print(f"ℹ️ lambda в патченных BitLinear-слоях: {lambdas}")
78 if lambdas and max(lambdas) < 1.0:
79 print(
80 "⚠️ lambda < 1.0 -- веса ЕЩЁ НЕ тернарные (QAT не завершён либо "
81 "не запускался). Экспортируется текущее состояние как есть -- "
82 "это ожидаемо для baseline sanity-теста из раздела 2 инструкции."
83 )
84
85 raw_sd = model.state_dict()
86 hf_sd = {}
87 dropped = []
88 for k, v in raw_sd.items():
89 if k.endswith("lambda_"):
90 dropped.append(k)
91 continue
92 hf_sd[k] = v.detach().to(dtype).contiguous()
93
94 print(f"🧹 Отброшено служебных ключей (BitLinear.lambda_): {len(dropped)}")
95 print(f"📦 Тензоров к сохранению: {len(hf_sd)}")
96
97 mtp_like = [k for k in hf_sd if k.startswith("mtp.")]
98 print(f"🔎 mtp.* тензоров в результате: {len(mtp_like)}")
99 if not mtp_like:
100 raise RuntimeError(
101 "В экспорте нет mtp.* тензоров -- JiRackDeltaNet_27b.py "
102 "не обновлён до версии с MTPBlock. Обнови файл на сервере.")
103
104 os.makedirs(args.out_dir, exist_ok=True)
105 out_path = os.path.join(args.out_dir, "model.safetensors")
106 save_file(hf_sd, out_path, metadata={"format": "pt"})
107 print(f"✅ Сохранено: {out_path}")
108
109 if args.hf_mtp_dir:
110 print("\n✅ MTP-голова в экспорте -- реальные веса из", args.hf_mtp_dir)
111 else:
112 print(
113 "\n⚠️ MTP-тензоры в экспорте есть, но если .pt был старый и "
114 "--hf_mtp_dir не задан -- это СЛУЧАЙНЫЕ веса. Для рабочего "
115 "спекулятивного декода перезапусти с --hf_mtp_dir."
116 )
117
118
119 if __name__ == "__main__":
120 main()
121