import torch import torch.onnx from JiRackTernaryPyTorch_1b_inf import TernaryTransformer1B, TernaryConfig def export_onnx(): device = torch.device("cpu") # Экспорт лучше делать на CPU для стабильности графа checkpoint_path = "/root/JiRackTernary1/new/model_packed.safetensors" onnx_output = "/root/JiRackTernary1/new/model/jirack_1b.onnx" # 1. Инициализируем модель config = TernaryConfig() model = TernaryTransformer1B(config).to(device) # 2. Загружаем веса (наш метод из наследника всё сделает сам) model.load_prod_weights(checkpoint_path, device) model.eval() # 3. Создаем "пустышку" для входа (dummy input) # Batch size 1, Sequence length 128 dummy_input = torch.randint(0, config.vocab_size, (1, 128)).to(device) print(f"🚀 Начинаю экспорт в ONNX...") torch.onnx.export( model, (dummy_input,), onnx_output, export_params=True, # Сохраняем веса внутри файла opset_version=17, # Используем свежий opset для поддержки продвинутых функций do_constant_folding=True, # Оптимизация графа (распаковка гаммы вмерджится в веса!) input_names=['input_ids'], output_names=['logits'], dynamic_axes={ 'input_ids': {0: 'batch_size', 1: 'seq_len'}, 'logits': {0: 'batch_size', 1: 'seq_len'} } ) print(f"✅ Готово! Модель сохранена: {onnx_output}") if __name__ == "__main__": export_onnx()