25 lines
999 B
Python
25 lines
999 B
Python
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}")
|