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