File size: 1,906 Bytes
9b92c75
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
from __future__ import annotations

import argparse
from pathlib import Path

import torch
from torch import nn

from .config import apply_overrides, load_config
from .model import build_model


class ExportModel(nn.Module):
    def __init__(self, model: nn.Module) -> None:
        super().__init__()
        self.model = model

    def forward(self, images: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]:
        outputs = self.model(images)
        return outputs["pred_logits"], outputs["pred_boxes"]


def main() -> None:
    parser = argparse.ArgumentParser(description="Export ObjectModel-v1 to ONNX")
    parser.add_argument("--config", default="configs/objectmodel_v1.yaml")
    parser.add_argument("--checkpoint", required=True)
    parser.add_argument("--output", default="objectmodel-v1.onnx")
    parser.add_argument("--opset", type=int, default=20)
    parser.add_argument("--set", action="append", default=[])
    args = parser.parse_args()
    config = apply_overrides(load_config(args.config), args.set)
    model = build_model(config)
    checkpoint = torch.load(args.checkpoint, map_location="cpu", weights_only=False)
    model.load_state_dict(checkpoint.get("ema", checkpoint.get("model", checkpoint)))
    model.eval()
    wrapper = ExportModel(model)
    size = model.spec.input_size
    sample = torch.randn(1, 3, size, size)
    output = Path(args.output)
    output.parent.mkdir(parents=True, exist_ok=True)
    torch.onnx.export(
        wrapper,
        (sample,),
        output,
        input_names=["images"],
        output_names=["logits", "boxes"],
        dynamic_axes={
            "images": {0: "batch"},
            "logits": {0: "batch"},
            "boxes": {0: "batch"},
        },
        opset_version=args.opset,
        dynamo=False,
    )
    print(f"Exported {output} ({output.stat().st_size / 1024**2:.2f} MiB)")


if __name__ == "__main__":
    main()