29 lines
785 B
Python
29 lines
785 B
Python
import os
|
|
import torch
|
|
from torchvision.models.segmentation import deeplabv3_resnet50
|
|
|
|
MODELO = "oak-1"
|
|
MODEL_NAME = "ervasModel"
|
|
NUM_CLASSES = 3
|
|
RESOLUCAO = (512, 512)
|
|
model_folder = os.path.join(MODELO, "backup")
|
|
model_name = MODEL_NAME + "_best"
|
|
|
|
# Carrega seu modelo (usa num_classes igual ao treino)
|
|
model = deeplabv3_resnet50(weights=None, num_classes=NUM_CLASSES)
|
|
model.load_state_dict(torch.load(os.path.join(model_folder, f"{model_name}.pth")))
|
|
model.eval()
|
|
|
|
# Dummy input com batch_size=1, channels=3, height=512, width=512
|
|
dummy_input = torch.randn(1, 3, RESOLUCAO[0], RESOLUCAO[1])
|
|
|
|
# Exporta
|
|
torch.onnx.export(
|
|
model,
|
|
dummy_input,
|
|
os.path.join(model_folder, f"{model_name}.onnx"),
|
|
input_names=["input"],
|
|
output_names=["output"],
|
|
opset_version=11
|
|
)
|