Image Segmentation
vision
image-detection
mathmanu commited on
Commit
f6b137e
Β·
verified Β·
1 Parent(s): 9e6bbdb

Add detr model files

Browse files
README.md ADDED
@@ -0,0 +1,195 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ ---
2
+ license: apache-2.0
3
+ tags:
4
+ - vision
5
+ - image-detection
6
+ - image-segmentation
7
+ datasets:
8
+ - COCO
9
+ ---
10
+
11
+ <div align="center">
12
+
13
+ # DETR for TI EdgeAI
14
+
15
+ ### Set-Prediction Object Detection and Panoptic Segmentation via Transformers
16
+
17
+ [![License](https://img.shields.io/badge/License-Apache%202.0-blue?style=for-the-badge)](https://opensource.org/licenses/Apache-2.0)
18
+ [![Framework](https://img.shields.io/badge/Framework-ONNX-orange?style=for-the-badge)](https://onnx.ai/)
19
+ [![Task](https://img.shields.io/badge/Task-Detection%20%7C%20Segmentation-green?style=for-the-badge)](https://github.com/TexasInstruments/edgeai)
20
+ [![Dataset](https://img.shields.io/badge/Dataset-COCO-blueviolet?style=for-the-badge)](https://cocodataset.org)
21
+
22
+ </div>
23
+
24
+ ---
25
+
26
+ ## Overview
27
+
28
+ **DETR** (DEtection TRansformer) is a transformer-based object detection architecture from Facebook Research that eliminates the need for hand-crafted components like anchor generation and NMS post-processing. It reformulates object detection as a direct set-prediction problem, using bipartite matching together with a transformer encoder-decoder to predict a fixed set of 100 object queries per image.
29
+
30
+ DETR matches Faster R-CNN with a ResNet-50 backbone in AP while using half the FLOPs, and extends naturally to panoptic segmentation by adding a mask head on top of the detection queries. The original [facebookresearch/detr](https://github.com/facebookresearch/detr) repository is archived (Apache 2.0), and pretrained COCO weights are pulled automatically via `torch.hub`.
31
+
32
+ This export covers the four ResNet-backbone **detection** variants (`detr_resnet50`, `detr_resnet50_dc5`, `detr_resnet101`, `detr_resnet101_dc5`). The panoptic segmentation variants (`detr_resnet50_panoptic`, `detr_resnet50_dc5_panoptic`, `detr_resnet101_panoptic`) are supported by the upstream repository and by `prepare_model.py`, but are not included as pre-exported artifacts in this folder.
33
+
34
+ ---
35
+
36
+ ## Model Variants
37
+
38
+ | Model | Backbone | mAP[.5:.95]% | mAP[.50]% | Validated Devices | Config |
39
+ |-------|----------|-------------|-----------|--------------------|--------|
40
+ | `detr_resnet50` | ResNet-50 | 42.0 | 62.4 | TDA4VH | [detr_resnet50_config.yaml](detr_resnet50_config.yaml) |
41
+ | `detr_resnet50_dc5` | ResNet-50 DC5 | 43.3 | 63.1 | TDA4VH | [detr_resnet50_dc5_config.yaml](detr_resnet50_dc5_config.yaml) |
42
+ | `detr_resnet101` | ResNet-101 | 43.5 | 63.8 | TDA4VH | [detr_resnet101_config.yaml](detr_resnet101_config.yaml) |
43
+ | `detr_resnet101_dc5` | ResNet-101 DC5 | 44.9 | 64.7 | TDA4VH | [detr_resnet101_dc5_config.yaml](detr_resnet101_dc5_config.yaml) |
44
+
45
+ > mAP values are on COCO val2017. DC5 = dilated convolutions in the last ResNet block (stride 16β†’32 kept at stride 8β†’16), giving higher-resolution features at the cost of higher compute.
46
+
47
+ **Recommended for edge deployment:** `detr_resnet50` (best accuracy/compute trade-off)
48
+
49
+ ---
50
+
51
+ ## Quick Start
52
+
53
+ ### Prerequisites
54
+
55
+ ```bash
56
+ pip install torch>=1.12.0 torchvision>=0.13.0 onnx>=1.14.0 scipy
57
+ pip install onnxruntime>=1.15.0
58
+ ```
59
+
60
+ `scipy` is required because DETR imports it at module load time (`models/matcher.py`). All of the above are auto-installed by `prepare_model.py` if missing.
61
+
62
+ ### Export the Model
63
+
64
+ ```bash
65
+ # Export the default model (detr_resnet50)
66
+ python prepare_model.py
67
+
68
+ # Export a specific model variant
69
+ python prepare_model.py --model detr_resnet101
70
+
71
+ # Export multiple variants at once
72
+ python prepare_model.py --model detr_resnet50 detr_resnet101
73
+
74
+ # List all available variants with accuracy info
75
+ python prepare_model.py --list-models
76
+
77
+ # Export from a locally trained checkpoint
78
+ python prepare_model.py --model detr_resnet50 --weights /path/to/checkpoint.pth
79
+ ```
80
+
81
+ The script automatically:
82
+ - Installs missing dependencies (`torch`, `torchvision`, `onnx`, `scipy`) if not present
83
+ - Loads the pretrained model via `torch.hub` (`facebookresearch/detr:main`), cloning the DETR source and downloading pretrained COCO weights from `dl.fbaipublicfiles.com` on first use
84
+ - Wraps the model to accept a plain `(N, 3, H, W)` tensor instead of a `NestedTensor`
85
+ - Exports to ONNX (opset 17 by default) with constant folding enabled, and validates the exported graph
86
+ - Saves the result as `<model_key>.onnx` in the output directory
87
+
88
+ > **Note:** DC5 (`_dc5`) variants are currently skipped by `prepare_model.py` with a warning, since TIDL does not yet support the dilated-conv backbone for compilation. The pre-exported `.onnx`/config artifacts for these variants remain in this folder for reference.
89
+
90
+ ### Compile and Infer uing edgeai-tidlrunner
91
+
92
+ > **Note:** Run the commands below from inside the `tidlrunner` directory (the cloned [edgeai-tidlrunner](https://github.com/TexasInstruments/edgeai-tidlrunner) repository), with `--config_path` pointing to this model's config file.
93
+
94
+ **Compile using edgeai-tidlrunner - on PC**
95
+
96
+ ```bash
97
+ cd /path/to/edgeai-tidlrunner
98
+ tidlrunner-cli compile --target_device J784S4 \
99
+ --config_path /path/to/detr_resnet50_config.yaml
100
+ ```
101
+
102
+ **Run Inference Benchmark - on device**
103
+
104
+ ```bash
105
+ cd /path/to/edgeai-tidlrunner
106
+ tidlrunner-cli infer --target_device J784S4 \
107
+ --config_path /path/to/detr_resnet50_config.yaml
108
+ ```
109
+
110
+ Swap `detr_resnet50_config.yaml` for `detr_resnet50_dc5_config.yaml`, `detr_resnet101_config.yaml`, or `detr_resnet101_dc5_config.yaml` to compile/infer the other variants.
111
+
112
+ ### Compile and Infer using edgeai-tidl-tools (Advanced):
113
+
114
+ Follow the instructions at https://github.com/TexasInstruments/edgeai-tidl-tools
115
+
116
+ ### Deploy using edgeai-tidl-tools:
117
+
118
+ Deplyment can be done using **[edgeai-tidl-tools](https://github.com/TexasInstruments/edgeai-tidl-tools)**. For ONNX models, onnxruntime-tidl with TIDL acceleration can be used. Consult the documentation of edgeai-tidl-tools for more details.
119
+
120
+ ---
121
+
122
+ ## Citation
123
+
124
+ If you use these models, please cite:
125
+
126
+ ```bibtex
127
+ @inproceedings{carion2020end,
128
+ title = {End-to-End Object Detection with Transformers},
129
+ author = {Carion, Nicolas and Massa, Francisco and Synnaeve, Gabriel and
130
+ Usunier, Nicolas and Kirillov, Alexander and Zagoruyko, Sergey},
131
+ booktitle = {European Conference on Computer Vision (ECCV)},
132
+ year = {2020}
133
+ }
134
+ ```
135
+
136
+ ---
137
+
138
+ ## πŸ”— Resources
139
+
140
+ | Resource | Link |
141
+ |----------|------|
142
+ | **Paper** | [arXiv:2005.12872](https://arxiv.org/abs/2005.12872) |
143
+ | **Source Code** | [facebookresearch/detr](https://github.com/facebookresearch/detr) |
144
+ | **Blog Post** | [End-to-End Object Detection with Transformers](https://ai.facebook.com/blog/end-to-end-object-detection-with-transformers) |
145
+ | **COCO Dataset** | [cocodataset.org](https://cocodataset.org) |
146
+ | **edgeai-tidl-tools** | [GitHub](https://github.com/TexasInstruments/edgeai-tidl-tools) |
147
+ | **edgeai-tidlrunner** | [GitHub](https://github.com/TexasInstruments/edgeai-tidlrunner) |
148
+ | **EdgeAI SDK** | [Documentation](https://github.com/TexasInstruments/edgeai/blob/main/edgeai-mpu/readme_sdk.md) |
149
+ | **TI EdgeAI Ecosystem** | [GitHub](https://github.com/TexasInstruments/edgeai) |
150
+
151
+ ---
152
+
153
+ ## Related Models
154
+
155
+ <table>
156
+ <tr>
157
+ <td align="center">
158
+
159
+ **Deformable-DETR**
160
+ Deformable attention
161
+ Faster convergence
162
+
163
+ </td>
164
+ <td align="center">
165
+
166
+ **RT-DETRv2**
167
+ Real-time transformer
168
+ NMS-free detection
169
+
170
+ </td>
171
+ <td align="center">
172
+
173
+ **RF-DETR**
174
+ Receptive-field DETR
175
+ Lightweight edge variant
176
+
177
+ </td>
178
+ <td align="center">
179
+
180
+ **DEIMv2**
181
+ Improved DETR training
182
+ Higher accuracy/epoch
183
+
184
+ </td>
185
+ </tr>
186
+ </table>
187
+
188
+ ---
189
+
190
+ <div align="center">
191
+
192
+ **Maintained by:** Texas Instruments EdgeAI Team
193
+ **Last Updated:** August 2026
194
+
195
+ </div>
detr_resnet101_config.yaml ADDED
@@ -0,0 +1,153 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ task_type: detection
2
+ dataloader:
3
+ name: coco_detection_dataloader
4
+ path: ./data/datasets/coco
5
+ preprocess:
6
+ resize: 800
7
+ crop: 800
8
+ data_layout: NCHW
9
+ reverse_channels: true
10
+ backend: cv2
11
+ interpolation: null
12
+ resize_with_pad: false
13
+ pad_color:
14
+ - 0
15
+ - 0
16
+ - 0
17
+ name: image_preprocess
18
+ session:
19
+ input_optimization: false
20
+ input_data_layout: NCHW
21
+ input_mean:
22
+ - 123.675
23
+ - 116.28
24
+ - 103.53
25
+ input_scale:
26
+ - 0.017125
27
+ - 0.017507
28
+ - 0.017429
29
+ runtime_options: {}
30
+ model_path: detr_resnet101.onnx
31
+ model_id: od-mh8047
32
+ input_details: null
33
+ output_details: null
34
+ num_inputs: 1
35
+ postprocess:
36
+ reshape_list: null
37
+ formatter:
38
+ name: DetectionXYWH2XYXYCenterXY
39
+ resize_with_pad: false
40
+ normalized_detections: true
41
+ shuffle_indices: null
42
+ squeeze_axis: null
43
+ ignore_index: null
44
+ model_output_type: split
45
+ logits_bbox_to_bbox_ls:
46
+ score_fn: sigmoid
47
+ bbox_index: 0
48
+ scores_index: 1
49
+ keypoint: false
50
+ object6dpose: false
51
+ name: detection_postprocess
52
+ metric:
53
+ label_offset_pred:
54
+ 1: 1
55
+ 2: 2
56
+ 3: 3
57
+ 4: 4
58
+ 5: 5
59
+ 6: 6
60
+ 7: 7
61
+ 8: 8
62
+ 9: 9
63
+ 10: 10
64
+ 11: 11
65
+ 12: 12
66
+ 13: 13
67
+ 14: 14
68
+ 15: 15
69
+ 16: 16
70
+ 17: 17
71
+ 18: 18
72
+ 19: 19
73
+ 20: 20
74
+ 21: 21
75
+ 22: 22
76
+ 23: 23
77
+ 24: 24
78
+ 25: 25
79
+ 26: 26
80
+ 27: 27
81
+ 28: 28
82
+ 29: 29
83
+ 30: 30
84
+ 31: 31
85
+ 32: 32
86
+ 33: 33
87
+ 34: 34
88
+ 35: 35
89
+ 36: 36
90
+ 37: 37
91
+ 38: 38
92
+ 39: 39
93
+ 40: 40
94
+ 41: 41
95
+ 42: 42
96
+ 43: 43
97
+ 44: 44
98
+ 45: 45
99
+ 46: 46
100
+ 47: 47
101
+ 48: 48
102
+ 49: 49
103
+ 50: 50
104
+ 51: 51
105
+ 52: 52
106
+ 53: 53
107
+ 54: 54
108
+ 55: 55
109
+ 56: 56
110
+ 57: 57
111
+ 58: 58
112
+ 59: 59
113
+ 60: 60
114
+ 61: 61
115
+ 62: 62
116
+ 63: 63
117
+ 64: 64
118
+ 65: 65
119
+ 66: 66
120
+ 67: 67
121
+ 68: 68
122
+ 69: 69
123
+ 70: 70
124
+ 71: 71
125
+ 72: 72
126
+ 73: 73
127
+ 74: 74
128
+ 75: 75
129
+ 76: 76
130
+ 77: 77
131
+ 78: 78
132
+ 79: 79
133
+ 80: 80
134
+ 81: 81
135
+ 82: 82
136
+ 83: 83
137
+ 84: 84
138
+ 85: 85
139
+ 86: 86
140
+ 87: 87
141
+ 88: 88
142
+ 89: 89
143
+ 90: 90
144
+ 91: 91
145
+ -1: -1
146
+ 0: 0
147
+ model_info:
148
+ metric_reference:
149
+ accuracy_ap[.5:.95]%: 43.5
150
+ model_shortlist: 10
151
+ compact_name: detr-r101-800x800
152
+ shortlisted: true
153
+ recommended: false
detr_resnet101_dc5_config.yaml ADDED
@@ -0,0 +1,153 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ task_type: detection
2
+ dataloader:
3
+ name: coco_detection_dataloader
4
+ path: ./data/datasets/coco
5
+ preprocess:
6
+ resize: 800
7
+ crop: 800
8
+ data_layout: NCHW
9
+ reverse_channels: true
10
+ backend: cv2
11
+ interpolation: null
12
+ resize_with_pad: false
13
+ pad_color:
14
+ - 0
15
+ - 0
16
+ - 0
17
+ name: image_preprocess
18
+ session:
19
+ input_optimization: false
20
+ input_data_layout: NCHW
21
+ input_mean:
22
+ - 123.675
23
+ - 116.28
24
+ - 103.53
25
+ input_scale:
26
+ - 0.017125
27
+ - 0.017507
28
+ - 0.017429
29
+ runtime_options: {}
30
+ model_path: detr_resnet101_dc5.onnx
31
+ model_id: od-mh8048
32
+ input_details: null
33
+ output_details: null
34
+ num_inputs: 1
35
+ postprocess:
36
+ reshape_list: null
37
+ formatter:
38
+ name: DetectionXYWH2XYXYCenterXY
39
+ resize_with_pad: false
40
+ normalized_detections: true
41
+ shuffle_indices: null
42
+ squeeze_axis: null
43
+ ignore_index: null
44
+ model_output_type: split
45
+ logits_bbox_to_bbox_ls:
46
+ score_fn: sigmoid
47
+ bbox_index: 0
48
+ scores_index: 1
49
+ keypoint: false
50
+ object6dpose: false
51
+ name: detection_postprocess
52
+ metric:
53
+ label_offset_pred:
54
+ 1: 1
55
+ 2: 2
56
+ 3: 3
57
+ 4: 4
58
+ 5: 5
59
+ 6: 6
60
+ 7: 7
61
+ 8: 8
62
+ 9: 9
63
+ 10: 10
64
+ 11: 11
65
+ 12: 12
66
+ 13: 13
67
+ 14: 14
68
+ 15: 15
69
+ 16: 16
70
+ 17: 17
71
+ 18: 18
72
+ 19: 19
73
+ 20: 20
74
+ 21: 21
75
+ 22: 22
76
+ 23: 23
77
+ 24: 24
78
+ 25: 25
79
+ 26: 26
80
+ 27: 27
81
+ 28: 28
82
+ 29: 29
83
+ 30: 30
84
+ 31: 31
85
+ 32: 32
86
+ 33: 33
87
+ 34: 34
88
+ 35: 35
89
+ 36: 36
90
+ 37: 37
91
+ 38: 38
92
+ 39: 39
93
+ 40: 40
94
+ 41: 41
95
+ 42: 42
96
+ 43: 43
97
+ 44: 44
98
+ 45: 45
99
+ 46: 46
100
+ 47: 47
101
+ 48: 48
102
+ 49: 49
103
+ 50: 50
104
+ 51: 51
105
+ 52: 52
106
+ 53: 53
107
+ 54: 54
108
+ 55: 55
109
+ 56: 56
110
+ 57: 57
111
+ 58: 58
112
+ 59: 59
113
+ 60: 60
114
+ 61: 61
115
+ 62: 62
116
+ 63: 63
117
+ 64: 64
118
+ 65: 65
119
+ 66: 66
120
+ 67: 67
121
+ 68: 68
122
+ 69: 69
123
+ 70: 70
124
+ 71: 71
125
+ 72: 72
126
+ 73: 73
127
+ 74: 74
128
+ 75: 75
129
+ 76: 76
130
+ 77: 77
131
+ 78: 78
132
+ 79: 79
133
+ 80: 80
134
+ 81: 81
135
+ 82: 82
136
+ 83: 83
137
+ 84: 84
138
+ 85: 85
139
+ 86: 86
140
+ 87: 87
141
+ 88: 88
142
+ 89: 89
143
+ 90: 90
144
+ 91: 91
145
+ -1: -1
146
+ 0: 0
147
+ model_info:
148
+ metric_reference:
149
+ accuracy_ap[.5:.95]%: 44.9
150
+ model_shortlist: 10
151
+ compact_name: detr-r101-dc5-800x800
152
+ shortlisted: true
153
+ recommended: false
detr_resnet50_config.yaml ADDED
@@ -0,0 +1,153 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ task_type: detection
2
+ dataloader:
3
+ name: coco_detection_dataloader
4
+ path: ./data/datasets/coco
5
+ preprocess:
6
+ resize: 800
7
+ crop: 800
8
+ data_layout: NCHW
9
+ reverse_channels: true
10
+ backend: cv2
11
+ interpolation: null
12
+ resize_with_pad: false
13
+ pad_color:
14
+ - 0
15
+ - 0
16
+ - 0
17
+ name: image_preprocess
18
+ session:
19
+ input_optimization: false
20
+ input_data_layout: NCHW
21
+ input_mean:
22
+ - 123.675
23
+ - 116.28
24
+ - 103.53
25
+ input_scale:
26
+ - 0.017125
27
+ - 0.017507
28
+ - 0.017429
29
+ runtime_options: {}
30
+ model_path: detr_resnet50.onnx
31
+ model_id: od-mh8045
32
+ input_details: null
33
+ output_details: null
34
+ num_inputs: 1
35
+ postprocess:
36
+ reshape_list: null
37
+ formatter:
38
+ name: DetectionXYWH2XYXYCenterXY
39
+ resize_with_pad: false
40
+ normalized_detections: true
41
+ shuffle_indices: null
42
+ squeeze_axis: null
43
+ ignore_index: null
44
+ model_output_type: split
45
+ logits_bbox_to_bbox_ls:
46
+ score_fn: sigmoid
47
+ bbox_index: 0
48
+ scores_index: 1
49
+ keypoint: false
50
+ object6dpose: false
51
+ name: detection_postprocess
52
+ metric:
53
+ label_offset_pred:
54
+ 1: 1
55
+ 2: 2
56
+ 3: 3
57
+ 4: 4
58
+ 5: 5
59
+ 6: 6
60
+ 7: 7
61
+ 8: 8
62
+ 9: 9
63
+ 10: 10
64
+ 11: 11
65
+ 12: 12
66
+ 13: 13
67
+ 14: 14
68
+ 15: 15
69
+ 16: 16
70
+ 17: 17
71
+ 18: 18
72
+ 19: 19
73
+ 20: 20
74
+ 21: 21
75
+ 22: 22
76
+ 23: 23
77
+ 24: 24
78
+ 25: 25
79
+ 26: 26
80
+ 27: 27
81
+ 28: 28
82
+ 29: 29
83
+ 30: 30
84
+ 31: 31
85
+ 32: 32
86
+ 33: 33
87
+ 34: 34
88
+ 35: 35
89
+ 36: 36
90
+ 37: 37
91
+ 38: 38
92
+ 39: 39
93
+ 40: 40
94
+ 41: 41
95
+ 42: 42
96
+ 43: 43
97
+ 44: 44
98
+ 45: 45
99
+ 46: 46
100
+ 47: 47
101
+ 48: 48
102
+ 49: 49
103
+ 50: 50
104
+ 51: 51
105
+ 52: 52
106
+ 53: 53
107
+ 54: 54
108
+ 55: 55
109
+ 56: 56
110
+ 57: 57
111
+ 58: 58
112
+ 59: 59
113
+ 60: 60
114
+ 61: 61
115
+ 62: 62
116
+ 63: 63
117
+ 64: 64
118
+ 65: 65
119
+ 66: 66
120
+ 67: 67
121
+ 68: 68
122
+ 69: 69
123
+ 70: 70
124
+ 71: 71
125
+ 72: 72
126
+ 73: 73
127
+ 74: 74
128
+ 75: 75
129
+ 76: 76
130
+ 77: 77
131
+ 78: 78
132
+ 79: 79
133
+ 80: 80
134
+ 81: 81
135
+ 82: 82
136
+ 83: 83
137
+ 84: 84
138
+ 85: 85
139
+ 86: 86
140
+ 87: 87
141
+ 88: 88
142
+ 89: 89
143
+ 90: 90
144
+ 91: 91
145
+ -1: -1
146
+ 0: 0
147
+ model_info:
148
+ metric_reference:
149
+ accuracy_ap[.5:.95]%: 42.0
150
+ model_shortlist: 10
151
+ compact_name: detr-r50-800x800
152
+ shortlisted: true
153
+ recommended: true
detr_resnet50_dc5_config.yaml ADDED
@@ -0,0 +1,153 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ task_type: detection
2
+ dataloader:
3
+ name: coco_detection_dataloader
4
+ path: ./data/datasets/coco
5
+ preprocess:
6
+ resize: 800
7
+ crop: 800
8
+ data_layout: NCHW
9
+ reverse_channels: true
10
+ backend: cv2
11
+ interpolation: null
12
+ resize_with_pad: false
13
+ pad_color:
14
+ - 0
15
+ - 0
16
+ - 0
17
+ name: image_preprocess
18
+ session:
19
+ input_optimization: false
20
+ input_data_layout: NCHW
21
+ input_mean:
22
+ - 123.675
23
+ - 116.28
24
+ - 103.53
25
+ input_scale:
26
+ - 0.017125
27
+ - 0.017507
28
+ - 0.017429
29
+ runtime_options: {}
30
+ model_path: detr_resnet50_dc5.onnx
31
+ model_id: od-mh8046
32
+ input_details: null
33
+ output_details: null
34
+ num_inputs: 1
35
+ postprocess:
36
+ reshape_list: null
37
+ formatter:
38
+ name: DetectionXYWH2XYXYCenterXY
39
+ resize_with_pad: false
40
+ normalized_detections: true
41
+ shuffle_indices: null
42
+ squeeze_axis: null
43
+ ignore_index: null
44
+ model_output_type: split
45
+ logits_bbox_to_bbox_ls:
46
+ score_fn: sigmoid
47
+ bbox_index: 0
48
+ scores_index: 1
49
+ keypoint: false
50
+ object6dpose: false
51
+ name: detection_postprocess
52
+ metric:
53
+ label_offset_pred:
54
+ 1: 1
55
+ 2: 2
56
+ 3: 3
57
+ 4: 4
58
+ 5: 5
59
+ 6: 6
60
+ 7: 7
61
+ 8: 8
62
+ 9: 9
63
+ 10: 10
64
+ 11: 11
65
+ 12: 12
66
+ 13: 13
67
+ 14: 14
68
+ 15: 15
69
+ 16: 16
70
+ 17: 17
71
+ 18: 18
72
+ 19: 19
73
+ 20: 20
74
+ 21: 21
75
+ 22: 22
76
+ 23: 23
77
+ 24: 24
78
+ 25: 25
79
+ 26: 26
80
+ 27: 27
81
+ 28: 28
82
+ 29: 29
83
+ 30: 30
84
+ 31: 31
85
+ 32: 32
86
+ 33: 33
87
+ 34: 34
88
+ 35: 35
89
+ 36: 36
90
+ 37: 37
91
+ 38: 38
92
+ 39: 39
93
+ 40: 40
94
+ 41: 41
95
+ 42: 42
96
+ 43: 43
97
+ 44: 44
98
+ 45: 45
99
+ 46: 46
100
+ 47: 47
101
+ 48: 48
102
+ 49: 49
103
+ 50: 50
104
+ 51: 51
105
+ 52: 52
106
+ 53: 53
107
+ 54: 54
108
+ 55: 55
109
+ 56: 56
110
+ 57: 57
111
+ 58: 58
112
+ 59: 59
113
+ 60: 60
114
+ 61: 61
115
+ 62: 62
116
+ 63: 63
117
+ 64: 64
118
+ 65: 65
119
+ 66: 66
120
+ 67: 67
121
+ 68: 68
122
+ 69: 69
123
+ 70: 70
124
+ 71: 71
125
+ 72: 72
126
+ 73: 73
127
+ 74: 74
128
+ 75: 75
129
+ 76: 76
130
+ 77: 77
131
+ 78: 78
132
+ 79: 79
133
+ 80: 80
134
+ 81: 81
135
+ 82: 82
136
+ 83: 83
137
+ 84: 84
138
+ 85: 85
139
+ 86: 86
140
+ 87: 87
141
+ 88: 88
142
+ 89: 89
143
+ 90: 90
144
+ 91: 91
145
+ -1: -1
146
+ 0: 0
147
+ model_info:
148
+ metric_reference:
149
+ accuracy_ap[.5:.95]%: 43.3
150
+ model_shortlist: 10
151
+ compact_name: detr-r50-dc5-800x800
152
+ shortlisted: true
153
+ recommended: false
prepare_model.py ADDED
@@ -0,0 +1,702 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """Script to export DETR pretrained ONNX model(s).
2
+
3
+ DETR (Detection TRansformer) from Facebook Research.
4
+ Models are loaded via torch.hub, which automatically clones the DETR repository
5
+ and downloads pretrained COCO weights from dl.fbaipublicfiles.com on first use.
6
+
7
+ Reference: https://github.com/facebookresearch/detr
8
+
9
+ Detection variants (Apache 2.0, COCO pretrained):
10
+ detr_resnet50 – 800Γ—800, ~41M params, AP50:95 42.0, AP50 62.4
11
+ detr_resnet50_dc5 – 800Γ—800, ~41M params, AP50:95 43.3, AP50 63.1
12
+ detr_resnet101 – 800Γ—800, ~60M params, AP50:95 43.5, AP50 63.8
13
+ detr_resnet101_dc5 – 800Γ—800, ~60M params, AP50:95 44.9, AP50 64.7
14
+
15
+ Panoptic segmentation variants (Apache 2.0, COCO pretrained):
16
+ detr_resnet50_panoptic – 800Γ—800, ~43M params, PQ 43.4 (box AP 38.8)
17
+ detr_resnet50_dc5_panoptic – 800Γ—800, ~43M params, PQ 44.6 (box AP 40.2)
18
+ detr_resnet101_panoptic – 800Γ—800, ~62M params, PQ 45.1 (box AP 40.1)
19
+
20
+ DC5 = dilated convolutions in ResNet's last block (stride 16β†’32 β†’ stride 8β†’16),
21
+ yielding higher-resolution feature maps at the cost of increased computation.
22
+
23
+ ONNX inputs/outputs:
24
+ Input : images – (N, 3, H, W) float32, ImageNet-normalized
25
+ Output : pred_boxes – (N, 100, 4) boxes in (cx, cy, w, h), normalized [0, 1]
26
+ pred_logits – (N, 100, 92) class logits (det) or (N, 100, 251) (panoptic)
27
+ pred_masks – (N, 100, H/4, W/4) panoptic mask logits (panoptic only)
28
+
29
+ Notes:
30
+ - DETR always outputs exactly 100 query slots per image.
31
+ - Post-processing: apply softmax over pred_logits and filter out slots where the
32
+ no-object class (index 91 for detection, 250 for panoptic) has the highest score.
33
+ - DETR trains with variable-size inputs (shorter-side 800, max 1333). For ONNX a
34
+ fixed square shape is used (default 800Γ—800). Any size works; 800px gives best AP.
35
+ - First run requires internet access to clone the DETR repo and download weights.
36
+
37
+ Usage:
38
+ python prepare_model.py
39
+ python prepare_model.py --model detr_resnet50
40
+ python prepare_model.py --model detr_resnet50 detr_resnet101
41
+ python prepare_model.py --model detr_resnet50 --shape 800 1333
42
+ python prepare_model.py --model detr_resnet50 --weights /path/to/checkpoint.pth
43
+ python prepare_model.py --model detr_resnet50 --opset 18 --output-dir ./exports
44
+ python prepare_model.py --list-models
45
+ """
46
+
47
+ from __future__ import annotations
48
+
49
+ import argparse
50
+ import importlib
51
+ import os
52
+ import subprocess
53
+ import sys
54
+
55
+
56
+ # ─────────────────────────────────────────────
57
+ # Model catalogue
58
+ # ─────────────────────────────────────────────
59
+
60
+ # Each entry: variant_key β†’ metadata dict
61
+ # hub_name : function name used with torch.hub.load
62
+ # task : "detection" or "panoptic"
63
+ # num_classes: 91 for detection (outputs 92 logits incl. no-object),
64
+ # 250 for panoptic (outputs 251 logits incl. no-object)
65
+ MODEL_CATALOG: dict[str, dict] = {
66
+ # ── Detection ─────────────────────────────────────────────────────────────
67
+ "detr_resnet50": {
68
+ "hub_name": "detr_resnet50",
69
+ "shape": (800, 800),
70
+ "params_m": 41.3,
71
+ "ap50_95": 42.0,
72
+ "ap50": 62.4,
73
+ "pq": None,
74
+ "latency_ms": 36.0,
75
+ "license": "Apache 2.0",
76
+ "task": "detection",
77
+ "num_classes": 91,
78
+ "backbone": "ResNet-50",
79
+ "pth_url": "https://dl.fbaipublicfiles.com/detr/detr-r50-e632da11.pth",
80
+ },
81
+ "detr_resnet50_dc5": {
82
+ "hub_name": "detr_resnet50_dc5",
83
+ "shape": (800, 800),
84
+ "params_m": 41.3,
85
+ "ap50_95": 43.3,
86
+ "ap50": 63.1,
87
+ "pq": None,
88
+ "latency_ms": 83.0,
89
+ "license": "Apache 2.0",
90
+ "task": "detection",
91
+ "num_classes": 91,
92
+ "backbone": "ResNet-50 DC5",
93
+ "pth_url": "https://dl.fbaipublicfiles.com/detr/detr-r50-dc5-f0fb7ef5.pth",
94
+ },
95
+ "detr_resnet101": {
96
+ "hub_name": "detr_resnet101",
97
+ "shape": (800, 800),
98
+ "params_m": 60.0,
99
+ "ap50_95": 43.5,
100
+ "ap50": 63.8,
101
+ "pq": None,
102
+ "latency_ms": 50.0,
103
+ "license": "Apache 2.0",
104
+ "task": "detection",
105
+ "num_classes": 91,
106
+ "backbone": "ResNet-101",
107
+ "pth_url": "https://dl.fbaipublicfiles.com/detr/detr-r101-2c7b67e5.pth",
108
+ },
109
+ "detr_resnet101_dc5": {
110
+ "hub_name": "detr_resnet101_dc5",
111
+ "shape": (800, 800),
112
+ "params_m": 60.0,
113
+ "ap50_95": 44.9,
114
+ "ap50": 64.7,
115
+ "pq": None,
116
+ "latency_ms": 97.0,
117
+ "license": "Apache 2.0",
118
+ "task": "detection",
119
+ "num_classes": 91,
120
+ "backbone": "ResNet-101 DC5",
121
+ "pth_url": "https://dl.fbaipublicfiles.com/detr/detr-r101-dc5-a2e86def.pth",
122
+ },
123
+ # ── Panoptic segmentation ──────────────────────────────────────────────────
124
+ "detr_resnet50_panoptic": {
125
+ "hub_name": "detr_resnet50_panoptic",
126
+ "shape": (800, 800),
127
+ "params_m": 43.2,
128
+ "ap50_95": 38.8,
129
+ "ap50": None,
130
+ "pq": 43.4,
131
+ "latency_ms": None,
132
+ "license": "Apache 2.0",
133
+ "task": "panoptic",
134
+ "num_classes": 250,
135
+ "backbone": "ResNet-50",
136
+ "pth_url": "https://dl.fbaipublicfiles.com/detr/detr-r50-panoptic-00ce5173.pth",
137
+ },
138
+ "detr_resnet50_dc5_panoptic": {
139
+ "hub_name": "detr_resnet50_dc5_panoptic",
140
+ "shape": (800, 800),
141
+ "params_m": 43.2,
142
+ "ap50_95": 40.2,
143
+ "ap50": None,
144
+ "pq": 44.6,
145
+ "latency_ms": None,
146
+ "license": "Apache 2.0",
147
+ "task": "panoptic",
148
+ "num_classes": 250,
149
+ "backbone": "ResNet-50 DC5",
150
+ "pth_url": "https://dl.fbaipublicfiles.com/detr/detr-r50-dc5-panoptic-da08f1b1.pth",
151
+ },
152
+ "detr_resnet101_panoptic": {
153
+ "hub_name": "detr_resnet101_panoptic",
154
+ "shape": (800, 800),
155
+ "params_m": 62.0,
156
+ "ap50_95": 40.1,
157
+ "ap50": None,
158
+ "pq": 45.1,
159
+ "latency_ms": None,
160
+ "license": "Apache 2.0",
161
+ "task": "panoptic",
162
+ "num_classes": 250,
163
+ "backbone": "ResNet-101",
164
+ "pth_url": "https://dl.fbaipublicfiles.com/detr/detr-r101-panoptic-40021d53.pth",
165
+ },
166
+ }
167
+
168
+ DEFAULT_MODEL = "detr_resnet50"
169
+
170
+ # torch.hub repo string for DETR
171
+ _HUB_REPO = "facebookresearch/detr:main"
172
+
173
+
174
+ # ─────────────────────────────────────────────
175
+ # Dependency management
176
+ # ─────────────────────────────────────────────
177
+
178
+ def _pip_install(*packages: str) -> None:
179
+ """Install *packages* via pip, suppressing verbose output."""
180
+ print(f"[DEP] Installing: {', '.join(packages)} …")
181
+ result = subprocess.run(
182
+ [sys.executable, "-m", "pip", "install", *packages],
183
+ stdout=subprocess.DEVNULL,
184
+ stderr=subprocess.PIPE,
185
+ text=True,
186
+ )
187
+ if result.returncode != 0:
188
+ print(f"[DEP] ERROR: pip install failed (exit code {result.returncode}).")
189
+ if result.stderr:
190
+ print(result.stderr.strip())
191
+ print("[DEP] Please install manually and re-run:")
192
+ print(f" pip install {' '.join(packages)}")
193
+ sys.exit(1)
194
+ print("[DEP] Installation complete.\n")
195
+
196
+
197
+ def ensure_dependencies() -> None:
198
+ """Ensure torch, torchvision, onnx, and scipy are importable.
199
+
200
+ scipy is required because DETR's model code imports it at module load
201
+ time (scipy.optimize.linear_sum_assignment in models/matcher.py).
202
+ """
203
+ required = [
204
+ ("torch", "torch>=1.12.0"),
205
+ ("torchvision", "torchvision>=0.13.0"),
206
+ ("onnx", "onnx>=1.14.0"),
207
+ ("scipy", "scipy"),
208
+ ]
209
+ missing_pip = []
210
+ for mod_name, pip_spec in required:
211
+ try:
212
+ importlib.import_module(mod_name)
213
+ print(f"[DEP] βœ” {mod_name} is installed.")
214
+ except ImportError:
215
+ print(f"[DEP] ✘ {mod_name} not found.")
216
+ missing_pip.append(pip_spec)
217
+
218
+ if missing_pip:
219
+ _pip_install(*missing_pip)
220
+
221
+ print()
222
+
223
+
224
+ # ─────────────────────────────────────────────
225
+ # Hub path helpers
226
+ # ─────────────────────────────────────────────
227
+
228
+ def _add_detr_to_path() -> str:
229
+ """Add the downloaded DETR source directory to sys.path (index 0).
230
+
231
+ torch.hub.load clones facebookresearch/detr to
232
+ ``<hub_dir>/facebookresearch_detr_main/``. This directory must be on
233
+ sys.path so that ``from util.misc import NestedTensor`` succeeds when
234
+ building the ONNX wrapper.
235
+
236
+ Returns the DETR root directory path.
237
+ """
238
+ import torch.hub as hub
239
+
240
+ hub_dir = hub.get_dir()
241
+ if not os.path.isdir(hub_dir):
242
+ raise RuntimeError(
243
+ f"torch.hub directory not found: {hub_dir}. "
244
+ "Run the script with internet access so torch.hub can clone DETR."
245
+ )
246
+
247
+ for entry in sorted(os.listdir(hub_dir), reverse=True):
248
+ if entry.startswith("facebookresearch_detr"):
249
+ detr_root = os.path.join(hub_dir, entry)
250
+ if os.path.isdir(detr_root):
251
+ if detr_root not in sys.path:
252
+ sys.path.insert(0, detr_root)
253
+ return detr_root
254
+
255
+ raise RuntimeError(
256
+ "Could not find DETR source in torch hub directory.\n"
257
+ f"Expected a subdirectory starting with 'facebookresearch_detr' inside {hub_dir}.\n"
258
+ "This is populated automatically by torch.hub.load on first use."
259
+ )
260
+
261
+
262
+ # ─────────────────────────────────────────────
263
+ # ONNX export wrappers
264
+ # ─────────────────────────────────────────────
265
+
266
+ def _make_wrapper(model, NestedTensor, task: str):
267
+ """Return an nn.Module that accepts a plain image tensor and produces flat outputs.
268
+
269
+ DETR's forward pass expects a NestedTensor (image + padding mask). These
270
+ wrappers create a zero mask (no padding) for fixed-size ONNX export, making
271
+ the model accept a standard (N, 3, H, W) float32 tensor.
272
+
273
+ Output order:
274
+ detection : pred_boxes (N,100,4), pred_logits (N,100,92)
275
+ panoptic : pred_boxes (N,100,4), pred_logits (N,100,251), pred_masks (N,100,H/4,W/4)
276
+ """
277
+ import torch
278
+ import torch.nn as nn
279
+
280
+ if task == "detection":
281
+ class _DetWrapper(nn.Module):
282
+ def __init__(self):
283
+ super().__init__()
284
+ self.model = model
285
+ self._NT = NestedTensor
286
+
287
+ def forward(self, images: torch.Tensor):
288
+ B, _, H, W = images.shape
289
+ mask = torch.zeros((B, H, W), dtype=torch.bool, device=images.device)
290
+ out = self.model(self._NT(images, mask))
291
+ return out["pred_boxes"], out["pred_logits"]
292
+
293
+ return _DetWrapper()
294
+
295
+ else: # panoptic
296
+ class _PanWrapper(nn.Module):
297
+ def __init__(self):
298
+ super().__init__()
299
+ self.model = model
300
+ self._NT = NestedTensor
301
+
302
+ def forward(self, images: torch.Tensor):
303
+ B, _, H, W = images.shape
304
+ mask = torch.zeros((B, H, W), dtype=torch.bool, device=images.device)
305
+ out = self.model(self._NT(images, mask))
306
+ return out["pred_boxes"], out["pred_logits"], out["pred_masks"]
307
+
308
+ return _PanWrapper()
309
+
310
+
311
+ # ─────────────────────────────────────────────
312
+ # Model catalogue helpers
313
+ # ─────────────────────────────────────────────
314
+
315
+ def print_model_table() -> None:
316
+ """Print a formatted table of all available models."""
317
+ col = 28
318
+ header = (
319
+ f" {'Variant':<{col}} {'Task':<10} {'Backbone':<16} "
320
+ f"{'Shape':<10} {'Params(M)':<10} {'AP50:95':<8} {'AP50/PQ':<8} "
321
+ f"{'Lat(ms)':<9} {'License'}"
322
+ )
323
+ sep = " " + "-" * (len(header) - 2)
324
+ print("\n" + "=" * len(header))
325
+ print(" Available DETR model variants")
326
+ print("=" * len(header))
327
+ print(header)
328
+ print(sep)
329
+
330
+ for key, info in MODEL_CATALOG.items():
331
+ h, w = info["shape"]
332
+ lat = f"{info['latency_ms']:.0f}" if info["latency_ms"] else "β€”"
333
+ ap50 = f"{info['ap50']:.1f}" if info["ap50"] is not None else f"PQ {info['pq']:.1f}"
334
+ print(
335
+ f" {key:<{col}} {info['task']:<10} {info['backbone']:<16} "
336
+ f"{h}Γ—{w:<5} {info['params_m']:<10.1f} {info['ap50_95']:<8.1f} "
337
+ f"{ap50:<8} {lat:<9} {info['license']}"
338
+ )
339
+ print("=" * len(header) + "\n")
340
+ print(" Latency measured on V100 GPU with TorchScript transformer.")
341
+ print(" DC5 = dilated conv in last ResNet block (higher-res features, slower).")
342
+ print(" AP values for detection on COCO val2017; PQ for panoptic on COCO val2017.\n")
343
+
344
+
345
+ # ─────────────────────────────────────────────
346
+ # Core export
347
+ # ─────────────────────────────────────────────
348
+
349
+ def export_model(
350
+ model_key: str,
351
+ output_dir: str,
352
+ shape: tuple[int, int] | None,
353
+ opset: int,
354
+ batch_size: int,
355
+ verbose: bool,
356
+ custom_weights: str | None,
357
+ force: bool,
358
+ force_hub_reload: bool,
359
+ ) -> str:
360
+ """Load a DETR model via torch.hub and export it to ONNX.
361
+
362
+ Pretrained COCO weights are downloaded automatically by torch.hub unless
363
+ *custom_weights* is provided.
364
+
365
+ Args:
366
+ model_key : Key from MODEL_CATALOG (e.g. "detr_resnet50").
367
+ output_dir : Final destination directory for the .onnx file.
368
+ shape : Custom (height, width) or None to use model default.
369
+ opset : ONNX opset version.
370
+ batch_size : Batch size embedded in the exported graph.
371
+ verbose : Show torch.hub download/loading messages.
372
+ custom_weights : Path to a local .pth checkpoint; None = COCO pretrained.
373
+ force : Re-export even if the destination .onnx already exists.
374
+ force_hub_reload: Force re-download of the DETR repo via torch.hub.
375
+
376
+ Returns:
377
+ Absolute path of the saved .onnx file.
378
+ """
379
+ import torch
380
+
381
+ info = MODEL_CATALOG[model_key]
382
+ hub_name = info["hub_name"]
383
+ task = info["task"]
384
+
385
+ # ── Resolve export shape ──────────────────────────────────────────────────
386
+ export_shape = shape if shape is not None else info["shape"]
387
+ h, w = export_shape
388
+
389
+ # ── Build destination path ────────────────────────────────────────────────
390
+ os.makedirs(output_dir, exist_ok=True)
391
+ shape_tag = f"_{h}x{w}" if shape is not None else ""
392
+ dst_name = f"{model_key}{shape_tag}.onnx"
393
+ dst_path = os.path.join(output_dir, dst_name)
394
+
395
+ if not force and os.path.exists(dst_path):
396
+ print(f"[SKIP] {dst_name} already exists. Use --force to re-export.\n")
397
+ return dst_path
398
+
399
+ print(f"[INFO] Model variant : {model_key}")
400
+ print(f"[INFO] Backbone : {info['backbone']}")
401
+ print(f"[INFO] Task : {task}")
402
+ print(f"[INFO] Input shape : {h}Γ—{w} (batch {batch_size})")
403
+ print(f"[INFO] ONNX opset : {opset}")
404
+ if custom_weights:
405
+ print(f"[INFO] Weights : {custom_weights}")
406
+ else:
407
+ print(f"[INFO] Weights : COCO pretrained (auto-downloaded)")
408
+ print(f"[INFO] Weight URL : {info['pth_url']}")
409
+ print()
410
+
411
+ # ── Load model via torch.hub ──────────────────────────────────────────────
412
+ print("[INFO] Loading model via torch.hub …")
413
+ print("[INFO] (First run will clone the DETR repo and download ~160–240 MB weights)")
414
+ if not verbose:
415
+ import warnings
416
+ warnings.filterwarnings("ignore")
417
+
418
+ load_kwargs: dict = {
419
+ "pretrained": custom_weights is None,
420
+ "force_reload": force_hub_reload,
421
+ }
422
+ try:
423
+ model = torch.hub.load(
424
+ _HUB_REPO, hub_name, trust_repo=True, verbose=verbose, **load_kwargs
425
+ )
426
+ except TypeError:
427
+ # PyTorch < 1.12 does not have trust_repo / verbose kwargs
428
+ model = torch.hub.load(_HUB_REPO, hub_name, **load_kwargs)
429
+
430
+ if custom_weights:
431
+ print(f"[INFO] Loading custom weights from: {custom_weights}")
432
+ checkpoint = torch.load(custom_weights, map_location="cpu")
433
+ state_dict = checkpoint.get("model", checkpoint)
434
+ model.load_state_dict(state_dict)
435
+
436
+ # Disable aux_loss to keep ONNX output clean (no aux_outputs in graph)
437
+ model.aux_loss = False
438
+ if hasattr(model, "detr"):
439
+ model.detr.aux_loss = False
440
+
441
+ model.eval()
442
+ print("[INFO] Model ready.\n")
443
+
444
+ # ── Import NestedTensor from DETR source ──────────────────────────────────
445
+ detr_root = _add_detr_to_path()
446
+ if verbose:
447
+ print(f"[INFO] DETR source : {detr_root}")
448
+ try:
449
+ from util.misc import NestedTensor # noqa: PLC0415
450
+ except ImportError as exc:
451
+ print(
452
+ f"[ERROR] Could not import NestedTensor from DETR source.\n"
453
+ f" Expected util/misc.py inside: {detr_root}\n"
454
+ f" Error: {exc}"
455
+ )
456
+ sys.exit(1)
457
+
458
+ # ── Build ONNX wrapper ────────────────────────────────────────────────────
459
+ wrapper = _make_wrapper(model, NestedTensor, task)
460
+ wrapper.eval()
461
+
462
+ # ── Dummy input ───────────────────────────────────────────────────────────
463
+ dummy = torch.zeros(batch_size, 3, h, w)
464
+
465
+ output_names = (
466
+ ["pred_boxes", "pred_logits", "pred_masks"]
467
+ if task == "panoptic"
468
+ else ["pred_boxes", "pred_logits"]
469
+ )
470
+
471
+ # ── Export ────────────────────────────────────────────────────────────────
472
+ print(f"[INFO] Exporting to ONNX (opset {opset}) …")
473
+ with torch.no_grad():
474
+ torch.onnx.export(
475
+ wrapper,
476
+ (dummy,),
477
+ dst_path,
478
+ input_names = ["images"],
479
+ output_names = output_names,
480
+ opset_version = opset,
481
+ do_constant_folding = True,
482
+ )
483
+
484
+ # ── Optional ONNX validation ──────────────────────────────────────────────
485
+ try:
486
+ import onnx # noqa: PLC0415
487
+ onnx_model = onnx.load(dst_path)
488
+ onnx.checker.check_model(onnx_model)
489
+ print("[INFO] ONNX model validation passed.")
490
+ except ImportError:
491
+ pass # onnx not available; skip validation
492
+ except Exception as exc:
493
+ print(f"[WARN] ONNX validation: {exc}")
494
+
495
+ size_mb = os.path.getsize(dst_path) / (1024 * 1024)
496
+ print(f"\n[SUCCESS] ONNX model saved to : {dst_path} ({size_mb:.1f} MB)\n")
497
+ return dst_path
498
+
499
+
500
+ # ─────────────────────────────────────────────
501
+ # CLI
502
+ # ─────────────────────────────────────────────
503
+
504
+ def build_parser() -> argparse.ArgumentParser:
505
+ default_output = os.path.dirname(os.path.abspath(__file__))
506
+
507
+ parser = argparse.ArgumentParser(
508
+ description=(
509
+ "Export DETR pretrained ONNX models.\n\n"
510
+ "Models are loaded via torch.hub (requires internet on first use).\n"
511
+ "Pretrained COCO weights are downloaded automatically from\n"
512
+ "dl.fbaipublicfiles.com. Run --list-models to see all variants."
513
+ ),
514
+ formatter_class=argparse.RawDescriptionHelpFormatter,
515
+ epilog=(
516
+ "Examples:\n"
517
+ " %(prog)s\n"
518
+ " %(prog)s --model detr_resnet50\n"
519
+ " %(prog)s --model detr_resnet50 detr_resnet101\n"
520
+ " %(prog)s --model detr_resnet50_dc5 detr_resnet101_dc5\n"
521
+ " %(prog)s --model detr_resnet50_panoptic detr_resnet101_panoptic\n"
522
+ " %(prog)s --model detr_resnet50 --shape 800 1333\n"
523
+ " %(prog)s --model detr_resnet50 --weights /path/to/checkpoint.pth\n"
524
+ " %(prog)s --model detr_resnet50 --opset 18 --output-dir ./exports\n"
525
+ " %(prog)s --list-models"
526
+ ),
527
+ )
528
+
529
+ # ── Model selection ───────────────────────────────────────────────────────
530
+ parser.add_argument(
531
+ "--model",
532
+ nargs="+",
533
+ default=[DEFAULT_MODEL],
534
+ choices=list(MODEL_CATALOG.keys()),
535
+ metavar="VARIANT",
536
+ help=(
537
+ f"Model variant(s) to export. Default: {DEFAULT_MODEL}. "
538
+ "Run --list-models to see all options."
539
+ ),
540
+ )
541
+
542
+ # ── Export parameters ─────────────────────────────────────────────────────
543
+ parser.add_argument(
544
+ "--shape",
545
+ nargs=2,
546
+ type=int,
547
+ default=None,
548
+ metavar=("H", "W"),
549
+ help=(
550
+ "Custom input resolution (height width). "
551
+ "DETR is flexible with input sizes; 800Γ—800 gives best accuracy. "
552
+ "Default: each model's native 800Γ—800."
553
+ ),
554
+ )
555
+ parser.add_argument(
556
+ "--opset",
557
+ type=int,
558
+ default=17,
559
+ metavar="N",
560
+ help="ONNX opset version. Default: 17.",
561
+ )
562
+ parser.add_argument(
563
+ "--batch-size",
564
+ type=int,
565
+ default=1,
566
+ metavar="N",
567
+ help="Batch size embedded in the exported ONNX graph. Default: 1.",
568
+ )
569
+
570
+ # ── Weight source ─────────────────────────────────────────────────────────
571
+ parser.add_argument(
572
+ "--weights",
573
+ default=None,
574
+ metavar="PATH",
575
+ help=(
576
+ "Path to a local .pth checkpoint (format: {'model': state_dict, ...}). "
577
+ "When omitted the official COCO pretrained weights are downloaded "
578
+ "automatically from dl.fbaipublicfiles.com via torch.hub."
579
+ ),
580
+ )
581
+
582
+ # ── Output ────────────────────────────────────────────────────────────────
583
+ parser.add_argument(
584
+ "--output-dir",
585
+ default=default_output,
586
+ metavar="DIR",
587
+ help=f"Directory where .onnx files will be saved. Default: {default_output}",
588
+ )
589
+ parser.add_argument(
590
+ "--force",
591
+ action="store_true",
592
+ default=False,
593
+ help="Re-export even if the destination .onnx file already exists.",
594
+ )
595
+
596
+ # ── Hub options ───────────────────────────────────────────────────────────
597
+ parser.add_argument(
598
+ "--force-hub-reload",
599
+ action="store_true",
600
+ default=False,
601
+ help=(
602
+ "Force torch.hub to re-clone the DETR repository and re-download "
603
+ "weights, bypassing the local cache. Use if the cache is corrupted."
604
+ ),
605
+ )
606
+
607
+ # ── Verbosity ─────────────────────────────────────────────────────────────
608
+ parser.add_argument(
609
+ "--quiet",
610
+ action="store_true",
611
+ default=False,
612
+ help="Suppress torch.hub download messages.",
613
+ )
614
+
615
+ # ── Utility ───────────────────────────────────────────────────────────────
616
+ parser.add_argument(
617
+ "--list-models",
618
+ action="store_true",
619
+ default=False,
620
+ help="Print the model catalogue table and exit.",
621
+ )
622
+
623
+ return parser
624
+
625
+
626
+ # ─────────────────────────────────────────────
627
+ # Entry point
628
+ # ─────────────────────────────────────────────
629
+
630
+ def main() -> None:
631
+ parser = build_parser()
632
+ args = parser.parse_args()
633
+
634
+ if args.list_models:
635
+ print_model_table()
636
+ return
637
+
638
+ # ── Warn when --weights is used with multiple models ─────────────────────
639
+ if args.weights and len(args.model) > 1:
640
+ print(
641
+ "[WARN] --weights applies the same checkpoint to every model in "
642
+ "--model.\n This is unusual; pass a single --model variant "
643
+ "when using custom weights."
644
+ )
645
+
646
+ # ── Install dependencies ──────────────────────────────────────────────────
647
+ ensure_dependencies()
648
+
649
+ # ── Export each model ─────────────────────────────────────────────────────
650
+ shape = (args.shape[0], args.shape[1]) if args.shape else None
651
+ output_dir = os.path.abspath(args.output_dir)
652
+
653
+ exported: list[str] = []
654
+ failed: list[str] = []
655
+
656
+ for model_key in args.model:
657
+ if '_dc5' in model_key:
658
+ print(f"[WARN] Model {model_key} is a DC5 variant and is temporarily disabled because TIDL does not support it. Skipping.")
659
+ continue
660
+
661
+ print(f"\n{'='*60}")
662
+ print(f" Exporting: {model_key}")
663
+ print(f"{'='*60}\n")
664
+
665
+ try:
666
+ out_path = export_model(
667
+ model_key = model_key,
668
+ output_dir = output_dir,
669
+ shape = shape,
670
+ opset = args.opset,
671
+ batch_size = args.batch_size,
672
+ verbose = not args.quiet,
673
+ custom_weights = args.weights,
674
+ force = args.force,
675
+ force_hub_reload = args.force_hub_reload,
676
+ )
677
+ exported.append(out_path)
678
+ except SystemExit:
679
+ raise
680
+ except Exception as exc:
681
+ print(f"[ERROR] Export failed for '{model_key}': {exc}")
682
+ failed.append(model_key)
683
+
684
+ # ── Summary ───────────────────────────────────────────────────────────────
685
+ print("\n" + "=" * 60)
686
+ print(" Export Summary")
687
+ print("=" * 60)
688
+ for path in exported:
689
+ size_mb = os.path.getsize(path) / (1024 * 1024)
690
+ print(f" βœ” {os.path.basename(path)} ({size_mb:.1f} MB)")
691
+ print(f" {path}")
692
+ if failed:
693
+ for key in failed:
694
+ print(f" ✘ {key} (FAILED)")
695
+ print("=" * 60 + "\n")
696
+
697
+ if failed:
698
+ sys.exit(1)
699
+
700
+
701
+ if __name__ == "__main__":
702
+ main()