1import torch
2import onnx
3import segmentation_models_pytorch as smp
4from torch.export.dynamic_shapes import Dim
5from onnxconverter_common import float16
6
7
8# Load model to plain smp.Unet
9path = "3_Class_FULL_FTW_Pretrained_singleWindow_v2.ckpt"
10
11ckpt = torch.load(path, map_location="cpu", weights_only=False)
12hparams = ckpt["hyper_parameters"]
13state_dict = {k.replace("model.", ""): v for k, v in ckpt["state_dict"].items()}
14del state_dict["criterion.weight"]
15print(hparams["model"], hparams["backbone"], hparams["in_channels"], hparams["num_classes"])
16model = smp.Unet(
17 encoder_name=hparams["backbone"],
18 encoder_weights=None,
19 in_channels=hparams["in_channels"],
20 classes=hparams["num_classes"],
21)
22model.eval()
23model.load_state_dict(state_dict, strict=True)
24
25# Export to exported program and then export to onnx
26program = torch.export.export(
27 model,
28 args=(torch.randn(1, hparams["in_channels"], 256, 256),),
29 dynamic_shapes={"x": (Dim.AUTO, hparams["in_channels"], Dim.AUTO, Dim.AUTO)},
30)
31output_path = "ftw-v2-single-window-unet-efficientnetb3-fp32.onnx"
32onnx_program = torch.onnx.export(model=program, f=output_path, external_data=False)
33onnx_model = onnx.load(output_path)
34onnx.checker.check_model(onnx_model)
35
36# Convert to fp16
37model_fp16 = float16.convert_float_to_float16(onnx_model, keep_io_types=True)
38onnx.save(model_fp16, output_path.replace("fp32", "fp16"))