import torch import os from safetensors.torch import save_file CHECKPOINT_PATH = "/mnt/nfs_share/JiRackTernaryPro_1b/jiarck_pro_1b_model.pt" OUTPUT_PATH = "/mnt/nfs_share/JiRackTernaryPro_1b/model.safetensors" def convert(): if not os.path.exists(CHECKPOINT_PATH): print(f"❌ Файл {CHECKPOINT_PATH} не найден!") return print(f"📦 Loading tensors from: {CHECKPOINT_PATH}") checkpoint = torch.load(CHECKPOINT_PATH, map_location="cpu", weights_only=False) # Достаем state_dict state_dict = checkpoint.get("model_state_dict", checkpoint) clean_sd = {} for k, v in state_dict.items(): if isinstance(v, torch.Tensor): # Убираем префиксы сразу при конвертации new_key = k.replace("_orig_mod.", "").replace("module.", "") # Safetensors требует contiguous тензоры. # Сохраняем оригинальный тип данных (int8/float32), чтобы не сломать perplexity. clean_sd[new_key] = v.detach().contiguous() os.makedirs(os.path.dirname(OUTPUT_PATH), exist_ok=True) print(f"💾 Saving {len(clean_sd)} tensors to Safetensors...") save_file(clean_sd, OUTPUT_PATH) size_gb = os.path.getsize(OUTPUT_PATH) / (1024**3) print(f"✅ Success! New size: {size_gb:.2f} GB") print(f"📍 Path: {OUTPUT_PATH}") if __name__ == "__main__": convert()