agrobot_base/Python/OAK/dlv3p_to_onnx.py

25 lines
999 B
Python
Raw Normal View History

2025-02-20 20:01:28 +00:00
import torch
import onnx
import network.modeling as models # 🔹 Usar a mesma biblioteca do treinamento!
# 🔹 Criar o modelo EXATO usado no treinamento
num_classes = 4
output_stride = 16
model = models.deeplabv3plus_resnet50(num_classes=num_classes, output_stride=output_stride) # 🔹 Backbone correto
# 🔹 Carregar os pesos salvos
checkpoint_path = "C:/AgroBaseModels/Ruas/model-1_0.pth"
checkpoint = torch.load(checkpoint_path, map_location=torch.device('cpu'))
model.load_state_dict(checkpoint) # 🔹 Agora os pesos vão carregar corretamente!
# 🔹 Colocar o modelo em modo de inferência
model.eval()
# 🔹 Criar entrada fictícia para exportação
dummy_input = torch.randn(1, 3, 512, 512) # 🔹 Ajustado para o tamanho usado no treinamento
# 🔹 Exportar para ONNX
onnx_path = "C:/AgroBaseModels/Ruas/model-1_0.onnx"
torch.onnx.export(model, dummy_input, onnx_path, opset_version=11, input_names=["input"], output_names=["output"])
print(f"Modelo salvo como {onnx_path}")