vision
image-detection
mathmanu commited on
Commit
a3a0f81
Β·
verified Β·
1 Parent(s): 9435a76

Add deformable_detr model files

Browse files
Files changed (3) hide show
  1. README.md +191 -0
  2. deformable_detr_single_scale_config.yaml +152 -0
  3. prepare_model.py +1292 -0
README.md ADDED
@@ -0,0 +1,191 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ ---
2
+ license: apache-2.0
3
+ tags:
4
+ - vision
5
+ - image-detection
6
+ datasets:
7
+ - COCO
8
+ ---
9
+
10
+ <div align="center">
11
+
12
+ # Deformable DETR for TI EdgeAI
13
+
14
+ ### Deformable Attention for Fast-Converging Transformer Detection
15
+
16
+ [![License](https://img.shields.io/badge/License-Apache%202.0-blue?style=for-the-badge)](https://opensource.org/licenses/Apache-2.0)
17
+ [![Framework](https://img.shields.io/badge/Framework-ONNX-orange?style=for-the-badge)](https://onnx.ai/)
18
+ [![Task](https://img.shields.io/badge/Task-Object%20Detection-green?style=for-the-badge)](https://github.com/TexasInstruments/edgeai)
19
+ [![Dataset](https://img.shields.io/badge/Dataset-COCO-blueviolet?style=for-the-badge)](https://cocodataset.org)
20
+
21
+ </div>
22
+
23
+ ---
24
+
25
+ ## Overview
26
+
27
+ **Deformable DETR** (Deformable Transformers for End-to-End Object Detection) is a transformer-based detector from SenseTime / fundamentalvision that addresses the slow convergence and limited feature resolution of the original DETR. Its key innovation is a **deformable attention module** that attends to only a small set of key sampling points around a reference point rather than all feature map positions, reducing complexity from **O(HΒ²WΒ²) to O(HW)**.
28
+
29
+ This efficient attention mechanism makes it practical to use **multi-scale feature maps**, which improves detection accuracy β€” especially on small objects β€” while training in **10Γ— fewer epochs** than DETR. Five variants are provided, ranging from a lightweight single-scale model to a two-stage design with iterative bounding box refinement.
30
+
31
+ Deformable DETR uses **300 query slots** (vs. 100 in DETR) and **sigmoid focal loss** for classification (no explicit background class); post-processing applies a score threshold rather than softmax + background filtering.
32
+
33
+ ---
34
+
35
+ ## Model Variants
36
+
37
+ | Model | Params | FLOPs | mAP[.5:.95]% | Validated Devices | Config |
38
+ |-------|--------|-------|--------------|--------------------|--------|
39
+ | `deformable_detr_single_scale` | 34M | 78G | 39.4 | TDA4VH | [deformable_detr_single_scale_config.yaml](deformable_detr_single_scale_config.yaml) |
40
+ | `deformable_detr_single_scale_dc5` | 34M | 128G | 41.5 | N/A | N/A |
41
+ | `deformable_detr` | 40M | 173G | 44.5 | N/A | N/A |
42
+ | `deformable_detr_plus_iterative_bbox_refinement` | 41M | 173G | 46.2 | N/A | N/A |
43
+ | `deformable_detr_two_stage` | 41M | 173G | 46.9 | N/A | N/A |
44
+
45
+ **Recommended for edge deployment:** `deformable_detr` (multi-scale, best accuracy/compute trade-off)
46
+
47
+ > mAP values on COCO val2017. All variants use a ResNet-50 backbone pretrained on ImageNet, 800Γ—800 input.
48
+ > The DC5 variant is disabled for TIDL deployment β€” TIDL does not support dilated convolutions in ResNet; its `.onnx` is provided for reference only.
49
+ > Only `deformable_detr_single_scale` currently ships with a validated TIDL config; the remaining variants have no `*_config.yaml` in this folder.
50
+
51
+ ---
52
+
53
+ ## Quick Start
54
+
55
+ ### Prerequisites
56
+
57
+ ```bash
58
+ # Core dependencies (auto-installed by prepare_model.py if missing)
59
+ pip install torch>=1.12.0 torchvision>=0.13.0 onnx>=1.14.0 scipy gdown>=5.2.0
60
+
61
+ # ONNX inference
62
+ pip install onnxruntime>=1.15.0
63
+ ```
64
+
65
+ ### Export the Model
66
+
67
+ ```bash
68
+ # List all available variants with accuracy and parameter info
69
+ python prepare_model.py --list-models
70
+
71
+ # Export the default model (deformable_detr - multi-scale, recommended)
72
+ python prepare_model.py
73
+
74
+ # Export a specific variant
75
+ python prepare_model.py --model deformable_detr_single_scale
76
+
77
+ # Export all variants (skips any already exported)
78
+ python prepare_model.py --model all
79
+
80
+ # Use HuggingFace Hub instead of Google Drive (recommended on corporate networks)
81
+ python prepare_model.py --method optimum --model all
82
+ ```
83
+
84
+ The script automatically:
85
+ - Installs missing dependencies (torch, torchvision, onnx, scipy, gdown) if not present
86
+ - Clones the [Deformable-DETR repository](https://github.com/fundamentalvision/Deformable-DETR) into `~/.cache/deformable_detr` (or downloads weights from HuggingFace Hub with `--method optimum`)
87
+ - Installs a pure-Python fallback for the multi-scale deformable attention module β€” no CUDA compilation required
88
+ - Downloads pretrained COCO weights and builds the model with the correct architecture flags
89
+ - Exports to ONNX (opset 17 by default), simplifies the graph with onnx-simplifier, and fixes float64 nodes for TIDL compatibility
90
+ - Validates the exported graph and saves it as `<model_key>.onnx`
91
+
92
+ ### Compile and Infer uing edgeai-tidlrunner
93
+
94
+ > **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.
95
+
96
+ **Compile using edgeai-tidlrunner - on PC**
97
+
98
+ ```bash
99
+ cd /path/to/edgeai-tidlrunner
100
+ tidlrunner-cli compile --target_device J784S4 \
101
+ --config_path /path/to/deformable_detr_single_scale_config.yaml
102
+ ```
103
+
104
+ **Run Inference Benchmark - on device**
105
+
106
+ ```bash
107
+ cd /path/to/edgeai-tidlrunner
108
+ tidlrunner-cli infer --target_device J784S4 \
109
+ --config_path /path/to/deformable_detr_single_scale_config.yaml
110
+ ```
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
+ ```bibtex
125
+ @article{zhu2020deformable,
126
+ title = {Deformable DETR: Deformable Transformers for End-to-End Object Detection},
127
+ author = {Zhu, Xizhou and Su, Weijie and Lu, Lewei and Li, Bin and
128
+ Wang, Xiaogang and Dai, Jifeng},
129
+ journal = {arXiv preprint arXiv:2010.04159},
130
+ year = {2020}
131
+ }
132
+ ```
133
+
134
+ ---
135
+
136
+ ## πŸ”— Resources
137
+
138
+ | Resource | Link |
139
+ |----------|------|
140
+ | **Paper** | [arXiv:2010.04159](https://arxiv.org/abs/2010.04159) |
141
+ | **Source Code** | [fundamentalvision/Deformable-DETR](https://github.com/fundamentalvision/Deformable-DETR) |
142
+ | **Dataset** | [COCO](https://cocodataset.org) |
143
+ | **edgeai-tidl-tools** | [GitHub](https://github.com/TexasInstruments/edgeai-tidl-tools) |
144
+ | **edgeai-tidlrunner** | [GitHub](https://github.com/TexasInstruments/edgeai-tidlrunner) |
145
+ | **EdgeAI SDK** | [Documentation](https://github.com/TexasInstruments/edgeai/blob/main/edgeai-mpu/readme_sdk.md) |
146
+
147
+ ---
148
+
149
+ ## Related Models
150
+
151
+ <table>
152
+ <tr>
153
+ <td align="center">
154
+
155
+ **DETR**
156
+ Original transformer detector
157
+ Predecessor to Deformable DETR
158
+
159
+ </td>
160
+ <td align="center">
161
+
162
+ **RT-DETRv2**
163
+ Real-time transformer detector
164
+ Modern DETR-style architecture
165
+
166
+ </td>
167
+ <td align="center">
168
+
169
+ **RF-DETR**
170
+ Receptive-field enhanced DETR
171
+ Recent DETR-family variant
172
+
173
+ </td>
174
+ <td align="center">
175
+
176
+ **DEIMv2**
177
+ Improved DETR training recipe
178
+ Faster convergence, higher accuracy
179
+
180
+ </td>
181
+ </tr>
182
+ </table>
183
+
184
+ ---
185
+
186
+ <div align="center">
187
+
188
+ **Maintained by:** Texas Instruments EdgeAI Team
189
+ **Last Updated:** August 2026
190
+
191
+ </div>
deformable_detr_single_scale_config.yaml ADDED
@@ -0,0 +1,152 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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: deformable_detr_single_scale.onnx
31
+ model_id: od-mh8060
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
+ 0: 0
146
+ model_info:
147
+ metric_reference:
148
+ accuracy_ap[.5:.95]%: 39.4
149
+ model_shortlist: 10
150
+ compact_name: deformable-detr-r50-ss-800x800
151
+ shortlisted: true
152
+ recommended: false
prepare_model.py ADDED
@@ -0,0 +1,1292 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """Script to export Deformable-DETR pretrained ONNX model(s).
2
+
3
+ Deformable DETR (Deformable Transformers for End-to-End Object Detection)
4
+ from SenseTime / fundamentalvision.
5
+
6
+ Reference: https://github.com/fundamentalvision/Deformable-DETR
7
+ Paper: https://arxiv.org/abs/2010.04159
8
+
9
+ Detection variants (Apache 2.0, COCO pretrained):
10
+ deformable_detr_single_scale – 800Γ—800, 34M params, AP50:95 39.4
11
+ deformable_detr_single_scale_dc5 – 800Γ—800, 34M params, AP50:95 41.5
12
+ deformable_detr – 800Γ—800, 40M params, AP50:95 44.5
13
+ deformable_detr_plus_iterative_bbox_refinement – 800Γ—800, 41M params, AP50:95 46.2
14
+ deformable_detr_two_stage – 800Γ—800, 41M params, AP50:95 46.9
15
+
16
+ Export methods (--method):
17
+ torch (default) – clones the official GitHub repo, downloads weights from
18
+ Google Drive via gdown. Requires internet access to
19
+ drive.google.com (may be blocked on corporate proxies).
20
+ optimum – downloads from HuggingFace Hub via the transformers
21
+ library. Proxy-friendly, no CUDA compilation, no Google
22
+ Drive. Uses HuggingFace model IDs under SenseTime/.
23
+
24
+ Notes:
25
+ - All variants use a ResNet-50 backbone, pre-trained on ImageNet.
26
+ - DC5 variant is disabled: TIDL does not support dilated convolution in ResNet.
27
+
28
+ Usage:
29
+ python prepare_model.py
30
+ python prepare_model.py --method optimum
31
+ python prepare_model.py --model deformable_detr_single_scale
32
+ python prepare_model.py --model deformable_detr_single_scale --method optimum
33
+ python prepare_model.py --model deformable_detr deformable_detr_two_stage
34
+ python prepare_model.py --model deformable_detr --shape 640 640
35
+ python prepare_model.py --model deformable_detr --weights /path/to/checkpoint.pth
36
+ python prepare_model.py --model deformable_detr --opset 18 --output-dir ./exports
37
+ python prepare_model.py --model deformable_detr --skip-simplify
38
+ python prepare_model.py --model all
39
+ python prepare_model.py --list-models
40
+ """
41
+
42
+ from __future__ import annotations
43
+
44
+ import argparse
45
+ import importlib
46
+ import math
47
+ import os
48
+ import subprocess
49
+ import sys
50
+
51
+
52
+ # ─────────────────────────────────────────────
53
+ # Model catalogue
54
+ # ─────────────────────────────────────────────
55
+
56
+ # Each entry: variant_key β†’ metadata dict
57
+ # num_feature_levels : 1 (single scale) or 4 (multi-scale)
58
+ # with_box_refine : iterative bounding box refinement
59
+ # two_stage : two-stage proposal + detection
60
+ # dilation : DC5 – dilation in ResNet's last block
61
+ # gdrive_id : Google Drive file ID (used by --method torch)
62
+ # hf_model_id : HuggingFace model ID (used by --method optimum)
63
+ MODEL_CATALOG: dict[str, dict] = {
64
+ "deformable_detr_single_scale": {
65
+ "num_feature_levels": 1,
66
+ "with_box_refine": False,
67
+ "two_stage": False,
68
+ "dilation": False,
69
+ "shape": (800, 800),
70
+ "params_m": 34,
71
+ "ap50_95": 39.4,
72
+ "flops_g": 78,
73
+ "fps_v100": 27.0,
74
+ "license": "Apache 2.0",
75
+ "gdrive_id": "1WEjQ9_FgfI5sw5OZZ4ix-OKk-IJ_-SDU",
76
+ "hf_model_id": "SenseTime/deformable-detr-single-scale",
77
+ },
78
+ "deformable_detr_single_scale_dc5": {
79
+ "num_feature_levels": 1,
80
+ "with_box_refine": False,
81
+ "two_stage": False,
82
+ "dilation": True,
83
+ "shape": (800, 800),
84
+ "params_m": 34,
85
+ "ap50_95": 41.5,
86
+ "flops_g": 128,
87
+ "fps_v100": 22.1,
88
+ "license": "Apache 2.0",
89
+ "gdrive_id": "1m_TgMjzH7D44fbA-c_jiBZ-xf-odxGdk",
90
+ "hf_model_id": "SenseTime/deformable-detr-single-scale-dc5",
91
+ },
92
+ "deformable_detr": {
93
+ "num_feature_levels": 4,
94
+ "with_box_refine": False,
95
+ "two_stage": False,
96
+ "dilation": False,
97
+ "shape": (800, 800),
98
+ "params_m": 40,
99
+ "ap50_95": 44.5,
100
+ "flops_g": 173,
101
+ "fps_v100": 15.0,
102
+ "license": "Apache 2.0",
103
+ "gdrive_id": "1nDWZWHuRwtwGden77NLM9JoWe-YisJnA",
104
+ "hf_model_id": "SenseTime/deformable-detr",
105
+ },
106
+ "deformable_detr_plus_iterative_bbox_refinement": {
107
+ "num_feature_levels": 4,
108
+ "with_box_refine": True,
109
+ "two_stage": False,
110
+ "dilation": False,
111
+ "shape": (800, 800),
112
+ "params_m": 41,
113
+ "ap50_95": 46.2,
114
+ "flops_g": 173,
115
+ "fps_v100": 15.0,
116
+ "license": "Apache 2.0",
117
+ "gdrive_id": "1JYKyRYzUH7uo9eVfDaVCiaIGZb5YTCuI",
118
+ "hf_model_id": "SenseTime/deformable-detr-with-box-refine",
119
+ },
120
+ "deformable_detr_two_stage": {
121
+ "num_feature_levels": 4,
122
+ "with_box_refine": True,
123
+ "two_stage": True,
124
+ "dilation": False,
125
+ "shape": (800, 800),
126
+ "params_m": 41,
127
+ "ap50_95": 46.9,
128
+ "flops_g": 173,
129
+ "fps_v100": 14.5,
130
+ "license": "Apache 2.0",
131
+ "gdrive_id": "15I03A7hNTpwuLNdfuEmW9_taZMNVssEp",
132
+ "hf_model_id": "SenseTime/deformable-detr-with-box-refine-two-stage",
133
+ },
134
+ }
135
+
136
+ DEFAULT_MODEL = "deformable_detr"
137
+
138
+ _REPO_URL = "https://github.com/fundamentalvision/Deformable-DETR.git"
139
+ _REPO_CACHE_DIR = os.path.join(os.path.expanduser("~"), ".cache", "deformable_detr")
140
+
141
+
142
+ # ─────────────────────────────────────────────
143
+ # Dependency installer
144
+ # ─────────────────────────────────────────────
145
+
146
+ def _pip_install(*packages: str) -> None:
147
+ """Install *packages* via pip, suppressing verbose output."""
148
+ print(f"[DEP] Installing: {', '.join(packages)} …")
149
+ result = subprocess.run(
150
+ [sys.executable, "-m", "pip", "install", *packages],
151
+ stdout=subprocess.DEVNULL,
152
+ stderr=subprocess.PIPE,
153
+ text=True,
154
+ )
155
+ if result.returncode != 0:
156
+ print(f"[DEP] ERROR: pip install failed (exit code {result.returncode}).")
157
+ if result.stderr:
158
+ print(result.stderr.strip())
159
+ print("[DEP] Please install manually and re-run:")
160
+ print(f" pip install {' '.join(packages)}")
161
+ sys.exit(1)
162
+ print("[DEP] Installation complete.\n")
163
+
164
+
165
+ def ensure_dependencies() -> None:
166
+ """Ensure torch, torchvision, onnx, scipy, gdown, and onnxsim are importable."""
167
+ required = [
168
+ ("torch", "torch>=1.12.0"),
169
+ ("torchvision", "torchvision>=0.13.0"),
170
+ ("onnx", "onnx>=1.14.0"),
171
+ ("onnxsim", "onnx-simplifier"),
172
+ ("scipy", "scipy"),
173
+ ("gdown", "gdown>=5.2.0"),
174
+ ]
175
+ missing = []
176
+ for mod_name, pip_spec in required:
177
+ try:
178
+ importlib.import_module(mod_name)
179
+ print(f"[DEP] βœ” {mod_name} is installed.")
180
+ except ImportError:
181
+ print(f"[DEP] ✘ {mod_name} not found.")
182
+ missing.append(pip_spec)
183
+
184
+ if missing:
185
+ _pip_install(*missing)
186
+ print()
187
+
188
+
189
+ def ensure_hf_dependencies() -> None:
190
+ """Ensure torch, torchvision, onnx, onnxsim, and transformers are importable
191
+ (needed for --method optimum)."""
192
+ required = [
193
+ ("torch", "torch>=1.12.0"),
194
+ ("torchvision", "torchvision>=0.13.0"),
195
+ ("onnx", "onnx>=1.14.0"),
196
+ ("onnxsim", "onnx-simplifier"),
197
+ ("transformers", "transformers>=4.30.0"),
198
+ ]
199
+ missing = []
200
+ for mod_name, pip_spec in required:
201
+ try:
202
+ importlib.import_module(mod_name)
203
+ print(f"[DEP] βœ” {mod_name} is installed.")
204
+ except ImportError:
205
+ print(f"[DEP] ✘ {mod_name} not found.")
206
+ missing.append(pip_spec)
207
+
208
+ if missing:
209
+ _pip_install(*missing)
210
+ print()
211
+
212
+
213
+ # ─────────────────────────────────────────────
214
+ # Repo setup
215
+ # ─────────────────────────────────────────────
216
+
217
+ def setup_repo(force_reclone: bool = False) -> str:
218
+ """Clone (or reuse) the Deformable-DETR repository.
219
+
220
+ Returns the absolute path to the repository root.
221
+ """
222
+ if os.path.isdir(_REPO_CACHE_DIR) and not force_reclone:
223
+ print(f"[REPO] Using cached repo: {_REPO_CACHE_DIR}")
224
+ return _REPO_CACHE_DIR
225
+
226
+ if os.path.isdir(_REPO_CACHE_DIR):
227
+ import shutil
228
+ shutil.rmtree(_REPO_CACHE_DIR)
229
+
230
+ os.makedirs(os.path.dirname(_REPO_CACHE_DIR), exist_ok=True)
231
+ print(f"[REPO] Cloning Deformable-DETR into {_REPO_CACHE_DIR} …")
232
+ result = subprocess.run(
233
+ ["git", "clone", "--depth", "1", _REPO_URL, _REPO_CACHE_DIR],
234
+ capture_output=True, text=True,
235
+ )
236
+ if result.returncode != 0:
237
+ print(f"[REPO] ERROR: git clone failed.\n{result.stderr.strip()}")
238
+ sys.exit(1)
239
+ print("[REPO] Clone complete.\n")
240
+ return _REPO_CACHE_DIR
241
+
242
+
243
+ # ─────────────────────────────────────────────
244
+ # Torchvision compatibility shim
245
+ # ─────────────────────────────────────────────
246
+
247
+ def _patch_torchvision_compat() -> None:
248
+ """Stub out removed torchvision symbols referenced by Deformable-DETR's
249
+ util/misc.py.
250
+
251
+ The repo does ``float(torchvision.__version__[:3]) < 0.5`` to gate
252
+ old code. For torchvision >= 0.10 the string ``"0.15"[:3]`` is
253
+ ``"0.1"`` β†’ float 0.1 β†’ the condition is True, triggering an import
254
+ of ``_NewEmptyTensorOp`` that was removed in torchvision 0.9.
255
+
256
+ We add a harmless stub so the import succeeds. The function
257
+ util/misc.interpolate() falls back to ``torch.nn.functional.interpolate``
258
+ for non-empty tensors (all practical cases), so the stub is never called.
259
+ """
260
+ import torch # noqa: PLC0415
261
+ import torchvision.ops.misc as _tvm # noqa: PLC0415
262
+
263
+ if hasattr(_tvm, "_NewEmptyTensorOp"):
264
+ return # already present (old torchvision) β€” nothing to do
265
+
266
+ class _NewEmptyTensorOp(torch.autograd.Function):
267
+ @staticmethod
268
+ def forward(ctx, x, new_size):
269
+ return x.new_empty(new_size)
270
+
271
+ @staticmethod
272
+ def backward(ctx, grad):
273
+ return grad, None
274
+
275
+ _tvm._NewEmptyTensorOp = _NewEmptyTensorOp
276
+ print("[COMPAT] Added _NewEmptyTensorOp stub to torchvision.ops.misc.\n")
277
+
278
+
279
+ # ─────────────────────────────────────────────
280
+ # Python fallback for deformable attention
281
+ # ─────────────────────────────────────────────
282
+
283
+ def _install_python_fallback(repo_root: str) -> None:
284
+ """Set up a pure-Python replacement for the multi-scale deformable
285
+ attention CUDA extension so that ONNX tracing works on CPU without
286
+ requiring CUDA compilation.
287
+
288
+ Strategy:
289
+ 1. Patch torchvision.ops.misc to add the removed _NewEmptyTensorOp stub
290
+ (required for util/misc.py to import on torchvision >= 0.10).
291
+ 2. Register a placeholder 'MultiScaleDeformableAttention' module in
292
+ sys.modules before any model code is imported (models/ops imports
293
+ this at module level).
294
+ 3. Import the pure-Python ms_deform_attn_core_pytorch function from
295
+ the repo source.
296
+ 4. Monkeypatch MSDeformAttn.forward to call ms_deform_attn_core_pytorch
297
+ directly, bypassing the MSDeformAttnFunction custom autograd op
298
+ (which has no ONNX symbolic and would break torch.onnx.export).
299
+ """
300
+ import types
301
+ import torch # noqa: PLC0415
302
+
303
+ if repo_root not in sys.path:
304
+ sys.path.insert(0, repo_root)
305
+
306
+ # Step 1 – Fix torchvision compatibility before importing any repo code.
307
+ _patch_torchvision_compat()
308
+
309
+ # Step 2 – Register a stub MSDA module so models/ops imports succeed.
310
+ if "MultiScaleDeformableAttention" not in sys.modules:
311
+ stub = types.ModuleType("MultiScaleDeformableAttention")
312
+ sys.modules["MultiScaleDeformableAttention"] = stub
313
+ print("[OPS] Registered stub MultiScaleDeformableAttention module.")
314
+
315
+ # Step 3 – Import the Python-only reference implementation.
316
+ from models.ops.functions.ms_deform_attn_func import ( # noqa: PLC0415
317
+ ms_deform_attn_core_pytorch,
318
+ )
319
+
320
+ # Step 4 – Patch MSDeformAttn.forward to use ms_deform_attn_core_pytorch
321
+ # directly instead of calling MSDeformAttnFunction.apply.
322
+ # This is a verbatim rewrite of the original forward with only the final
323
+ # output = MSDeformAttnFunction.apply(...) line replaced.
324
+ import torch.nn.functional as F # noqa: PLC0415
325
+ import models.ops.modules.ms_deform_attn as _attn_mod # noqa: PLC0415
326
+
327
+ def _py_forward(
328
+ self,
329
+ query,
330
+ reference_points,
331
+ input_flatten,
332
+ input_spatial_shapes,
333
+ input_level_start_index,
334
+ input_padding_mask=None,
335
+ ):
336
+ N, Len_q, _ = query.shape
337
+ N, Len_in, _ = input_flatten.shape
338
+ assert (input_spatial_shapes[:, 0] * input_spatial_shapes[:, 1]).sum() == Len_in
339
+
340
+ value = self.value_proj(input_flatten)
341
+ if input_padding_mask is not None:
342
+ value = value.masked_fill(input_padding_mask[..., None], float(0))
343
+ value = value.view(N, Len_in, self.n_heads, self.d_model // self.n_heads)
344
+
345
+ sampling_offsets = self.sampling_offsets(query).view(
346
+ N, Len_q, self.n_heads, self.n_levels, self.n_points, 2
347
+ )
348
+ attention_weights = self.attention_weights(query).view(
349
+ N, Len_q, self.n_heads, self.n_levels * self.n_points
350
+ )
351
+ attention_weights = F.softmax(attention_weights, -1).view(
352
+ N, Len_q, self.n_heads, self.n_levels, self.n_points
353
+ )
354
+
355
+ if reference_points.shape[-1] == 2:
356
+ offset_normalizer = torch.stack(
357
+ [input_spatial_shapes[..., 1], input_spatial_shapes[..., 0]], -1
358
+ )
359
+ sampling_locations = (
360
+ reference_points[:, :, None, :, None, :]
361
+ + sampling_offsets
362
+ / offset_normalizer[None, None, None, :, None, :]
363
+ )
364
+ elif reference_points.shape[-1] == 4:
365
+ sampling_locations = (
366
+ reference_points[:, :, None, :, None, :2]
367
+ + sampling_offsets
368
+ / self.n_points
369
+ * reference_points[:, :, None, :, None, 2:]
370
+ * 0.5
371
+ )
372
+ else:
373
+ raise ValueError(
374
+ f"Last dim of reference_points must be 2 or 4, "
375
+ f"got {reference_points.shape[-1]}"
376
+ )
377
+
378
+ output = ms_deform_attn_core_pytorch(
379
+ value, input_spatial_shapes, sampling_locations, attention_weights
380
+ )
381
+ output = self.output_proj(output)
382
+ return output
383
+
384
+ _attn_mod.MSDeformAttn.forward = _py_forward
385
+ print("[OPS] Pure-Python fallback installed for MSDeformAttn (no CUDA required).\n")
386
+
387
+
388
+ # ─────────────────────────────────────────────
389
+ # ONNX simplification
390
+ # ─────────────────────────────────────────────
391
+
392
+ def simplify_onnx(src_path: str, force: bool = False) -> bool:
393
+ """Run onnx-simplifier on *src_path* in-place.
394
+
395
+ Simplification folds constants, removes dead nodes, and cleans up
396
+ redundant ops produced by torch.onnx.export, making the graph smaller
397
+ and easier to deploy.
398
+
399
+ Falls back gracefully (copies as-is) if onnxsim is not installed.
400
+
401
+ Args:
402
+ src_path: Path to the .onnx file to simplify (modified in-place).
403
+ force : Re-run even if the file was already simplified.
404
+
405
+ Returns:
406
+ True on success (or if simplification was skipped gracefully).
407
+ """
408
+ print(f"\n[SIM] Running onnxsim on {os.path.basename(src_path)} …")
409
+
410
+ try:
411
+ import onnx # noqa: PLC0415
412
+ import onnxsim # noqa: PLC0415
413
+ except ImportError as exc:
414
+ missing = str(exc).split("'")[1] if "'" in str(exc) else str(exc)
415
+ print(f" [WARN] {missing} not installed – skipping simplification.")
416
+ print(" Install with: pip install onnx-simplifier")
417
+ return True
418
+
419
+ try:
420
+ model = onnx.load(src_path)
421
+ except Exception as exc:
422
+ print(f" [ERROR] Failed to load {src_path}: {exc}")
423
+ return False
424
+
425
+ try:
426
+ model_sim, check = onnxsim.simplify(model)
427
+ except Exception as exc:
428
+ print(f" [WARN] onnxsim failed: {exc} – keeping unsimplified model.")
429
+ return True
430
+
431
+ if not check:
432
+ print(" [WARN] onnxsim validation failed – keeping unsimplified model.")
433
+ return True
434
+
435
+ orig_nodes = len(model.graph.node)
436
+ sim_nodes = len(model_sim.graph.node)
437
+ delta = orig_nodes - sim_nodes
438
+ print(f" Nodes: {orig_nodes} β†’ {sim_nodes} (βˆ’{delta})")
439
+
440
+ try:
441
+ onnx.save(model_sim, src_path)
442
+ except Exception as exc:
443
+ print(f" [ERROR] Failed to save simplified model: {exc}")
444
+ return False
445
+
446
+ size_mb = os.path.getsize(src_path) / (1024 * 1024)
447
+ print(f"[OK] Simplified model saved: {src_path} ({size_mb:.1f} MB)")
448
+ return True
449
+
450
+
451
+ # ─────────────────────────────────────────────
452
+ # Weight download
453
+ # ─────────────────────────────────────────────
454
+
455
+ def download_weights(gdrive_id: str, weights_path: str, force: bool = False) -> bool:
456
+ """Download a pretrained checkpoint from Google Drive using gdown.
457
+
458
+ Args:
459
+ gdrive_id : Google Drive file ID.
460
+ weights_path: Local destination path for the .pth checkpoint.
461
+ force : Re-download even if file already exists.
462
+
463
+ Returns:
464
+ True on success.
465
+ """
466
+ import gdown # noqa: PLC0415
467
+
468
+ if os.path.exists(weights_path) and not force:
469
+ size_mb = os.path.getsize(weights_path) / 1024 / 1024
470
+ print(f"[SKIP] Weights already exist ({size_mb:.1f} MB). "
471
+ "Use --force to re-download.\n")
472
+ return True
473
+
474
+ url = f"https://drive.google.com/uc?id={gdrive_id}"
475
+ print(f"[DOWN] Downloading pretrained weights from Google Drive …")
476
+ print(f" File ID : {gdrive_id}")
477
+ print(f" Dest : {weights_path}")
478
+
479
+ # Pick up proxy settings from the environment (e.g. TI corporate proxy).
480
+ proxy = (
481
+ os.environ.get("HTTPS_PROXY")
482
+ or os.environ.get("https_proxy")
483
+ or os.environ.get("HTTP_PROXY")
484
+ or os.environ.get("http_proxy")
485
+ )
486
+ if proxy:
487
+ print(f" Proxy : {proxy}")
488
+
489
+ try:
490
+ dl_kwargs: dict = {"quiet": False}
491
+ if proxy:
492
+ dl_kwargs["proxy"] = proxy
493
+ gdown.download(url, weights_path, **dl_kwargs)
494
+ except Exception as exc:
495
+ print(f"[ERROR] gdown download failed: {exc}")
496
+ _print_manual_download_hint(gdrive_id, weights_path)
497
+ return False
498
+
499
+ if not os.path.exists(weights_path):
500
+ print("[ERROR] Download finished but file was not created.")
501
+ _print_manual_download_hint(gdrive_id, weights_path)
502
+ return False
503
+
504
+ size_mb = os.path.getsize(weights_path) / 1024 / 1024
505
+ print(f"[OK] Weights saved: {weights_path} ({size_mb:.1f} MB)\n")
506
+ return True
507
+
508
+
509
+ def _print_manual_download_hint(gdrive_id: str, weights_path: str) -> None:
510
+ """Print instructions for manually downloading a Google Drive checkpoint."""
511
+ url = f"https://drive.google.com/uc?id={gdrive_id}"
512
+ print(
513
+ f"\n[HINT] If you are behind a corporate proxy, download the checkpoint\n"
514
+ f" manually using one of the following commands:\n"
515
+ f"\n"
516
+ f" # with gdown and explicit proxy:\n"
517
+ f" gdown --proxy <proxy_url> '{url}' -O '{weights_path}'\n"
518
+ f"\n"
519
+ f" # with curl:\n"
520
+ f" curl -L -x <proxy_url> '{url}' -o '{weights_path}'\n"
521
+ f"\n"
522
+ f" Then re-run with --weights to skip the download:\n"
523
+ f" python prepare_model.py --model <variant> --weights '{weights_path}'\n"
524
+ )
525
+
526
+
527
+ # ─────────────────────────────────────────────
528
+ # Model catalogue helpers
529
+ # ─────────────────────────────────────────────
530
+
531
+ def print_model_table() -> None:
532
+ """Print a formatted table of all available model variants."""
533
+ col = 48
534
+ header = (
535
+ f" {'Variant':<{col}} {'Shape':<10} {'Params(M)':<10} "
536
+ f"{'AP50:95':<8} {'FLOPs(G)':<9} {'FPS(V100)':<10} {'License'}"
537
+ )
538
+ sep = " " + "-" * (len(header) - 2)
539
+ print("\n" + "=" * len(header))
540
+ print(" Available Deformable-DETR model variants")
541
+ print("=" * len(header))
542
+ print(header)
543
+ print(sep)
544
+
545
+ for key, info in MODEL_CATALOG.items():
546
+ h, w = info["shape"]
547
+ print(
548
+ f" {key:<{col}} {h}Γ—{w:<5} "
549
+ f"{info['params_m']:<10} {info['ap50_95']:<8.1f} "
550
+ f"{info['flops_g']:<9} {info['fps_v100']:<10.1f} "
551
+ f"{info['license']}"
552
+ )
553
+ print("=" * len(header) + "\n")
554
+ print(" All variants use ResNet-50 backbone, trained on COCO 2017.")
555
+ print(" AP50:95 measured on COCO val2017, inference speed on V100 GPU.\n")
556
+
557
+
558
+ # ─────────────────────────────────────────────
559
+ # Float64 removal
560
+ # ─────────────────────────────────────────────
561
+
562
+ def fix_float64_nodes(src_path: str) -> bool:
563
+ """Remove Cast-to-DOUBLE nodes and fix float64 initializers/constants
564
+ so the model is compatible with TIDL (which does not support float64).
565
+
566
+ torch.onnx.export inserts Cast(to=DOUBLE) nodes when Python-level float
567
+ literals (e.g. math.pi, which is float64) appear in position-encoding
568
+ computations. These nodes propagate float64 through most of the graph.
569
+
570
+ Strategy:
571
+ 1. Find every Cast node with to=DOUBLE.
572
+ 2. Re-wire each consumer of the Cast's output to use the Cast's input
573
+ (the upstream float32 tensor) directly, then delete the Cast node.
574
+ 3. Convert any float64 graph initializers to float32.
575
+ 4. Fix any Constant/ConstantOfShape attribute tensors that are DOUBLE.
576
+ 5. Validate with onnx.checker and save in-place.
577
+
578
+ Args:
579
+ src_path: Path to the .onnx file to fix (modified in-place).
580
+
581
+ Returns:
582
+ True on success or if no float64 tensors were found.
583
+ """
584
+ import numpy as np # noqa: PLC0415
585
+ try:
586
+ import onnx # noqa: PLC0415
587
+ from onnx import TensorProto, numpy_helper # noqa: PLC0415
588
+ except ImportError:
589
+ print(" [WARN] onnx not installed – skipping float64 fix.")
590
+ return True
591
+
592
+ try:
593
+ model = onnx.load(src_path)
594
+ except Exception as exc:
595
+ print(f" [ERROR] Failed to load {src_path}: {exc}")
596
+ return False
597
+
598
+ graph = model.graph
599
+
600
+ # Step 1 – Remove Cast-to-DOUBLE by re-wiring consumers to use Cast input.
601
+ consumers: dict = {}
602
+ for node in graph.node:
603
+ for inp in node.input:
604
+ consumers.setdefault(inp, []).append(node)
605
+
606
+ removed = 0
607
+ for node in list(graph.node):
608
+ if node.op_type != "Cast":
609
+ continue
610
+ for attr in node.attribute:
611
+ if attr.name == "to" and attr.i == TensorProto.DOUBLE:
612
+ cast_in = node.input[0]
613
+ cast_out = node.output[0]
614
+ for consumer in consumers.get(cast_out, []):
615
+ consumer.input[:] = [
616
+ cast_in if t == cast_out else t
617
+ for t in consumer.input
618
+ ]
619
+ graph.node.remove(node)
620
+ removed += 1
621
+ break
622
+
623
+ # Step 2 – Convert float64 graph initializers to float32.
624
+ init_fixed = 0
625
+ for init in graph.initializer:
626
+ if init.data_type == TensorProto.DOUBLE:
627
+ arr = numpy_helper.to_array(init).astype(np.float32)
628
+ init.CopyFrom(numpy_helper.from_array(arr, name=init.name))
629
+ init_fixed += 1
630
+
631
+ # Step 3 – Fix Constant/ConstantOfShape nodes with float64 value tensors.
632
+ # Use attr.type == TENSOR (the correct API) instead of attr.HasField("t"),
633
+ # which is unreliable across protobuf versions.
634
+ const_fixed = 0
635
+ for node in graph.node:
636
+ for attr in node.attribute:
637
+ if (attr.type == onnx.AttributeProto.TENSOR
638
+ and attr.t.data_type == TensorProto.DOUBLE):
639
+ arr = numpy_helper.to_array(attr.t).astype(np.float32)
640
+ attr.t.CopyFrom(numpy_helper.from_array(arr))
641
+ const_fixed += 1
642
+
643
+ # Step 4 – Update stale float64 type annotations in value_info.
644
+ # onnxsim stores intermediate tensor types in graph.value_info. When Cast-
645
+ # to-DOUBLE nodes are removed the stored annotations become stale and still
646
+ # say float64, which causes type-inference errors in TIDL and ONNX tools
647
+ # even though the actual computation is now float32.
648
+ vi_fixed = 0
649
+ for vi in list(graph.value_info) + list(graph.input) + list(graph.output):
650
+ if (vi.type.HasField("tensor_type")
651
+ and vi.type.tensor_type.elem_type == TensorProto.DOUBLE):
652
+ vi.type.tensor_type.elem_type = TensorProto.FLOAT
653
+ vi_fixed += 1
654
+
655
+ print(
656
+ f"[F64] Cast-to-DOUBLE removed: {removed}, "
657
+ f"initializers fixed: {init_fixed}, constants fixed: {const_fixed}, "
658
+ f"type annotations fixed: {vi_fixed}"
659
+ )
660
+
661
+ if removed == 0 and init_fixed == 0 and const_fixed == 0 and vi_fixed == 0:
662
+ print("[F64] No float64 tensors found – model already clean.")
663
+ return True
664
+
665
+ try:
666
+ onnx.checker.check_model(model)
667
+ print("[F64] ONNX model validation passed after float64 fix.")
668
+ except Exception as exc:
669
+ print(f"[WARN] ONNX validation after float64 fix: {exc}")
670
+
671
+ try:
672
+ onnx.save(model, src_path)
673
+ except Exception as exc:
674
+ print(f" [ERROR] Failed to save fixed model: {exc}")
675
+ return False
676
+
677
+ size_mb = os.path.getsize(src_path) / (1024 * 1024)
678
+ print(f"[OK] Float64-free model saved: {src_path} ({size_mb:.1f} MB)")
679
+ return True
680
+
681
+
682
+ # ─────────────────────────────────────────────
683
+ # Core export
684
+ # ─────────────────────────────────────────────
685
+
686
+ def export_model(
687
+ model_key: str,
688
+ output_dir: str,
689
+ shape: tuple[int, int] | None,
690
+ opset: int,
691
+ batch_size: int,
692
+ verbose: bool,
693
+ custom_weights: str | None,
694
+ force: bool,
695
+ force_reclone: bool,
696
+ skip_simplify: bool = False,
697
+ ) -> str:
698
+ """Clone the Deformable-DETR repo, download weights, and export to ONNX.
699
+
700
+ Args:
701
+ model_key : Key from MODEL_CATALOG.
702
+ output_dir : Directory to save the .onnx and .pth files.
703
+ shape : Custom (height, width) or None for model default.
704
+ opset : ONNX opset version.
705
+ batch_size : Batch size embedded in the exported graph.
706
+ verbose : Show additional progress messages.
707
+ custom_weights: Path to a local .pth checkpoint; None = pretrained.
708
+ force : Re-export even if .onnx already exists.
709
+ force_reclone : Force re-clone of the source repo.
710
+
711
+ Returns:
712
+ Absolute path of the saved .onnx file.
713
+ """
714
+ import torch # noqa: PLC0415
715
+
716
+ info = MODEL_CATALOG[model_key]
717
+ export_shape = shape if shape is not None else info["shape"]
718
+ h, w = export_shape
719
+
720
+ os.makedirs(output_dir, exist_ok=True)
721
+ shape_tag = f"_{h}x{w}" if shape is not None else ""
722
+ dst_name = f"{model_key}{shape_tag}.onnx"
723
+ dst_path = os.path.join(output_dir, dst_name)
724
+
725
+ if not force and os.path.exists(dst_path):
726
+ print(f"[SKIP] {dst_name} already exists. Use --force to re-export.\n")
727
+ return dst_path
728
+
729
+ print(f"[INFO] Variant : {model_key}")
730
+ print(f"[INFO] Feature lvls : {info['num_feature_levels']}")
731
+ print(f"[INFO] Box refine : {info['with_box_refine']}")
732
+ print(f"[INFO] Two-stage : {info['two_stage']}")
733
+ print(f"[INFO] DC5 dilation : {info['dilation']}")
734
+ print(f"[INFO] Input shape : {h}Γ—{w} (batch {batch_size})")
735
+ print(f"[INFO] ONNX opset : {opset}")
736
+ if custom_weights:
737
+ print(f"[INFO] Weights : {custom_weights}")
738
+ else:
739
+ print(f"[INFO] Weights : COCO pretrained (Google Drive)")
740
+
741
+ # ── Step 1: Clone repo ────────────────────────────────────────────────────
742
+ repo_root = setup_repo(force_reclone=force_reclone)
743
+
744
+ # ── Step 2: Install Python fallback for deformable attention ──────────────
745
+ _install_python_fallback(repo_root)
746
+
747
+ # ── Step 3: Download or locate weights ────────────────────────────────────
748
+ if custom_weights:
749
+ weights_path = custom_weights
750
+ if not os.path.exists(weights_path):
751
+ print(f"[ERROR] Custom weights not found: {weights_path}")
752
+ sys.exit(1)
753
+ else:
754
+ weights_path = os.path.join(output_dir, f"{model_key}.pth")
755
+ if not download_weights(info["gdrive_id"], weights_path, force=force):
756
+ print(f"[ERROR] Failed to download weights for '{model_key}'.")
757
+ print(
758
+ "\n[TIP] The default export method (torch) downloads weights from\n"
759
+ " Google Drive, which may be unreachable on corporate networks.\n"
760
+ " Try the HuggingFace-based method instead β€” no Google Drive\n"
761
+ " required, proxy-friendly:\n"
762
+ f"\n"
763
+ f" python prepare_model.py --method optimum --model {model_key}\n"
764
+ )
765
+ sys.exit(1)
766
+
767
+ # ── Step 4: Build model ───────────────────────────────────────────────────
768
+ print("[INFO] Building model …")
769
+ if repo_root not in sys.path:
770
+ sys.path.insert(0, repo_root)
771
+
772
+ from models import build_model # noqa: PLC0415
773
+
774
+ args = argparse.Namespace(
775
+ # Backbone
776
+ backbone = "resnet50",
777
+ dilation = info["dilation"],
778
+ position_embedding = "sine",
779
+ position_embedding_scale= 2 * math.pi,
780
+ num_feature_levels = info["num_feature_levels"],
781
+ # Transformer
782
+ enc_layers = 6,
783
+ dec_layers = 6,
784
+ dim_feedforward = 1024,
785
+ hidden_dim = 256,
786
+ dropout = 0.1,
787
+ nheads = 8,
788
+ num_queries = 300,
789
+ dec_n_points = 4,
790
+ enc_n_points = 4,
791
+ # Variant flags
792
+ with_box_refine = info["with_box_refine"],
793
+ two_stage = info["two_stage"],
794
+ # Segmentation (not used for detection export)
795
+ masks = False,
796
+ frozen_weights = None,
797
+ # Loss (needed by SetCriterion constructor, not used for inference)
798
+ aux_loss = False,
799
+ set_cost_class = 2.0,
800
+ set_cost_bbox = 5.0,
801
+ set_cost_giou = 2.0,
802
+ mask_loss_coef = 1.0,
803
+ dice_loss_coef = 1.0,
804
+ cls_loss_coef = 2.0,
805
+ bbox_loss_coef = 5.0,
806
+ giou_loss_coef = 2.0,
807
+ focal_alpha = 0.25,
808
+ # Dataset (determines num_classes = 91 for coco)
809
+ dataset_file = "coco",
810
+ coco_path = "./data/coco",
811
+ coco_panoptic_path = None,
812
+ remove_difficult = False,
813
+ # Device
814
+ device = "cpu",
815
+ )
816
+
817
+ model, _criterion, _postprocessors = build_model(args)
818
+ model.eval()
819
+ print("[INFO] Model built.\n")
820
+
821
+ # ── Step 5: Load pretrained weights ───────────────────────────────────────
822
+ print(f"[INFO] Loading weights from: {weights_path}")
823
+ checkpoint = torch.load(weights_path, map_location="cpu")
824
+ state_dict = checkpoint.get("model", checkpoint)
825
+ missing, unexpected = model.load_state_dict(state_dict, strict=False)
826
+ unexpected = [k for k in unexpected if not k.endswith(("total_params", "total_ops"))]
827
+ if missing:
828
+ print(f"[WARN] Missing keys : {missing[:5]}{'…' if len(missing) > 5 else ''}")
829
+ if unexpected:
830
+ print(f"[WARN] Unexpected keys: {unexpected[:5]}{'…' if len(unexpected) > 5 else ''}")
831
+ print("[INFO] Weights loaded.\n")
832
+
833
+ # ── Step 6: Build ONNX wrapper ────────────────────────────────────────────
834
+ import torch # noqa: PLC0415
835
+ import torch.nn as nn # noqa: PLC0415
836
+ from util.misc import NestedTensor # noqa: PLC0415
837
+
838
+ class _Wrapper(nn.Module):
839
+ def __init__(self):
840
+ super().__init__()
841
+ self.model = model
842
+ self._NT = NestedTensor
843
+
844
+ def forward(self, images: torch.Tensor):
845
+ B, _, H, W = images.shape
846
+ mask = torch.zeros((B, H, W), dtype=torch.bool, device=images.device)
847
+ out = self.model(self._NT(images, mask))
848
+ return out["pred_boxes"], out["pred_logits"]
849
+
850
+ wrapper = _Wrapper().eval()
851
+
852
+ # ── Step 7: Export to ONNX ────────────────────────────────────────────────
853
+ print(f"[INFO] Exporting to ONNX (opset {opset}) …")
854
+ dummy = torch.zeros(batch_size, 3, h, w)
855
+
856
+ with torch.no_grad():
857
+ torch.onnx.export(
858
+ wrapper,
859
+ (dummy,),
860
+ dst_path,
861
+ input_names = ["images"],
862
+ output_names = ["pred_boxes", "pred_logits"],
863
+ opset_version = opset,
864
+ do_constant_folding = True,
865
+ )
866
+
867
+ # ── Optional ONNX validation ──────────────────────────────────────────────
868
+ try:
869
+ import onnx # noqa: PLC0415
870
+ onnx_model = onnx.load(dst_path)
871
+ onnx.checker.check_model(onnx_model)
872
+ print("[INFO] ONNX model validation passed.")
873
+ except ImportError:
874
+ pass
875
+ except Exception as exc:
876
+ print(f"[WARN] ONNX validation: {exc}")
877
+
878
+ # ── Optional onnxsim simplification ──────────────────────────────────────
879
+ if not skip_simplify:
880
+ simplify_onnx(dst_path, force=force)
881
+
882
+ # ── Fix float64 nodes for TIDL compatibility ──────────────────────────────
883
+ fix_float64_nodes(dst_path)
884
+
885
+ size_mb = os.path.getsize(dst_path) / (1024 * 1024)
886
+ print(f"\n[SUCCESS] ONNX model saved to: {dst_path} ({size_mb:.1f} MB)\n")
887
+ return dst_path
888
+
889
+
890
+ # ─────────────────────────────────────────────
891
+ # Optimum / HuggingFace export
892
+ # ─────────────────────────────────────────────
893
+
894
+ def export_model_optimum(
895
+ model_key: str,
896
+ output_dir: str,
897
+ shape: tuple[int, int] | None,
898
+ opset: int,
899
+ batch_size: int,
900
+ force: bool,
901
+ skip_simplify: bool = False,
902
+ ) -> str:
903
+ """Download from HuggingFace and export Deformable-DETR to ONNX.
904
+
905
+ Uses the HuggingFace transformers implementation of Deformable DETR,
906
+ which is a pure-Python port of the original architecture. No Google
907
+ Drive access, no CUDA compilation, and no repo cloning required.
908
+
909
+ The transformers model is downloaded via HuggingFace Hub. Proxy
910
+ settings are picked up automatically from the HTTPS_PROXY / https_proxy
911
+ environment variables (TI corporate proxy is supported).
912
+
913
+ Args:
914
+ model_key : Key from MODEL_CATALOG.
915
+ output_dir : Directory to save the .onnx file.
916
+ shape : Custom (height, width) or None for model default.
917
+ opset : ONNX opset version.
918
+ batch_size : Batch size embedded in the exported graph.
919
+ force : Re-export even if .onnx already exists.
920
+
921
+ Returns:
922
+ Absolute path of the saved .onnx file.
923
+ """
924
+ import torch # noqa: PLC0415
925
+ import torch.nn as nn # noqa: PLC0415
926
+ from transformers import DeformableDetrForObjectDetection # noqa: PLC0415
927
+
928
+ info = MODEL_CATALOG[model_key]
929
+ hf_model_id = info["hf_model_id"]
930
+ export_shape = shape if shape is not None else info["shape"]
931
+ h, w = export_shape
932
+
933
+ os.makedirs(output_dir, exist_ok=True)
934
+ shape_tag = f"_{h}x{w}" if shape is not None else ""
935
+ dst_name = f"{model_key}{shape_tag}.onnx"
936
+ dst_path = os.path.join(output_dir, dst_name)
937
+
938
+ if not force and os.path.exists(dst_path):
939
+ print(f"[SKIP] {dst_name} already exists. Use --force to re-export.\n")
940
+ return dst_path
941
+
942
+ print(f"[INFO] Method : optimum (HuggingFace transformers)")
943
+ print(f"[INFO] HF model ID : {hf_model_id}")
944
+ print(f"[INFO] Input shape : {h}Γ—{w} (batch {batch_size})")
945
+ print(f"[INFO] ONNX opset : {opset}")
946
+ print()
947
+
948
+ # ── Download / load from HuggingFace ─────────────────────────────────────
949
+ print(f"[INFO] Loading model from HuggingFace …")
950
+ print("[INFO] (First run downloads ~150–200 MB; cached at ~/.cache/huggingface/)")
951
+ model = DeformableDetrForObjectDetection.from_pretrained(hf_model_id)
952
+ model.eval()
953
+ print("[INFO] Model ready.\n")
954
+
955
+ # ── ONNX export wrapper ───────────────────────────────────────────────────
956
+ # DeformableDetrForObjectDetection.forward(pixel_values, pixel_mask=None)
957
+ # outputs: DeformableDetrObjectDetectionOutput with .pred_boxes and .logits
958
+ # We rename logits β†’ pred_logits to match our postprocess configs.
959
+ class _HFWrapper(nn.Module):
960
+ def __init__(self):
961
+ super().__init__()
962
+ self.model = model
963
+
964
+ def forward(self, images: torch.Tensor):
965
+ out = self.model(pixel_values=images)
966
+ return out.pred_boxes, out.logits
967
+
968
+ wrapper = _HFWrapper().eval()
969
+ dummy = torch.zeros(batch_size, 3, h, w)
970
+
971
+ # ── Export ────────────────────────────────────────────────────────────────
972
+ print(f"[INFO] Exporting to ONNX (opset {opset}) …")
973
+ with torch.no_grad():
974
+ torch.onnx.export(
975
+ wrapper,
976
+ (dummy,),
977
+ dst_path,
978
+ input_names = ["images"],
979
+ output_names = ["pred_boxes", "pred_logits"],
980
+ opset_version = opset,
981
+ do_constant_folding= True,
982
+ )
983
+
984
+ # ── Optional validation ───────────────────────────────────────────────────
985
+ try:
986
+ import onnx # noqa: PLC0415
987
+ onnx_model = onnx.load(dst_path)
988
+ onnx.checker.check_model(onnx_model)
989
+ print("[INFO] ONNX model validation passed.")
990
+ except ImportError:
991
+ pass
992
+ except Exception as exc:
993
+ print(f"[WARN] ONNX validation: {exc}")
994
+
995
+ # ── Optional onnxsim simplification ──────────────────────────────────────
996
+ if not skip_simplify:
997
+ simplify_onnx(dst_path, force=force)
998
+
999
+ # ── Fix float64 nodes for TIDL compatibility ──────────────────────────────
1000
+ fix_float64_nodes(dst_path)
1001
+
1002
+ size_mb = os.path.getsize(dst_path) / (1024 * 1024)
1003
+ print(f"\n[SUCCESS] ONNX model saved to: {dst_path} ({size_mb:.1f} MB)\n")
1004
+ return dst_path
1005
+
1006
+
1007
+ # ─────────────────────────────────────────────
1008
+ # CLI
1009
+ # ─────────────────────────────────────────────
1010
+
1011
+ def build_parser() -> argparse.ArgumentParser:
1012
+ default_output = os.path.dirname(os.path.abspath(__file__))
1013
+
1014
+ parser = argparse.ArgumentParser(
1015
+ description=(
1016
+ "Export Deformable-DETR pretrained ONNX models.\n\n"
1017
+ "The Deformable-DETR source is cloned from GitHub on first use.\n"
1018
+ "Pretrained COCO weights are downloaded from Google Drive via gdown.\n"
1019
+ "Run --list-models to see all available variants."
1020
+ ),
1021
+ formatter_class=argparse.RawDescriptionHelpFormatter,
1022
+ epilog=(
1023
+ "Examples:\n"
1024
+ " %(prog)s\n"
1025
+ " %(prog)s --method optimum # proxy-friendly HF download\n"
1026
+ " %(prog)s --model deformable_detr_single_scale\n"
1027
+ " %(prog)s --model deformable_detr_single_scale --method optimum\n"
1028
+ " %(prog)s --model deformable_detr deformable_detr_two_stage\n"
1029
+ " %(prog)s --model deformable_detr --shape 640 640\n"
1030
+ " %(prog)s --model deformable_detr --weights /path/to/checkpoint.pth\n"
1031
+ " %(prog)s --model deformable_detr --opset 18 --output-dir ./exports\n"
1032
+ " %(prog)s --model deformable_detr --skip-simplify\n"
1033
+ " %(prog)s --model all\n"
1034
+ " %(prog)s --list-models"
1035
+ ),
1036
+ )
1037
+
1038
+ parser.add_argument(
1039
+ "--model",
1040
+ nargs="+",
1041
+ default=[DEFAULT_MODEL],
1042
+ choices=list(MODEL_CATALOG.keys()) + ["all"],
1043
+ metavar="VARIANT",
1044
+ help=(
1045
+ f"Model variant(s) to export. Default: {DEFAULT_MODEL}. "
1046
+ "Use 'all' to export every variant that has not yet been exported. "
1047
+ "Run --list-models to see all options."
1048
+ ),
1049
+ )
1050
+ parser.add_argument(
1051
+ "--shape",
1052
+ nargs=2,
1053
+ type=int,
1054
+ default=None,
1055
+ metavar=("H", "W"),
1056
+ help=(
1057
+ "Custom input resolution (height width). "
1058
+ "Default: each model's native 800Γ—800."
1059
+ ),
1060
+ )
1061
+ parser.add_argument(
1062
+ "--opset",
1063
+ type=int,
1064
+ default=17,
1065
+ metavar="N",
1066
+ help="ONNX opset version. Default: 17.",
1067
+ )
1068
+ parser.add_argument(
1069
+ "--batch-size",
1070
+ type=int,
1071
+ default=1,
1072
+ metavar="N",
1073
+ help="Batch size embedded in the exported ONNX graph. Default: 1.",
1074
+ )
1075
+ parser.add_argument(
1076
+ "--weights",
1077
+ default=None,
1078
+ metavar="PATH",
1079
+ help=(
1080
+ "Path to a local .pth checkpoint (format: {'model': state_dict, ...}). "
1081
+ "When omitted the official COCO pretrained weights are downloaded "
1082
+ "automatically from Google Drive."
1083
+ ),
1084
+ )
1085
+ parser.add_argument(
1086
+ "--output-dir",
1087
+ default=default_output,
1088
+ metavar="DIR",
1089
+ help=f"Directory where .onnx and .pth files will be saved. Default: {default_output}",
1090
+ )
1091
+ parser.add_argument(
1092
+ "--force",
1093
+ action="store_true",
1094
+ default=False,
1095
+ help="Re-export and re-download even if output files already exist.",
1096
+ )
1097
+ parser.add_argument(
1098
+ "--force-reclone",
1099
+ action="store_true",
1100
+ default=False,
1101
+ help=(
1102
+ "Force re-clone of the Deformable-DETR repository, "
1103
+ "removing the cached copy in ~/.cache/deformable_detr."
1104
+ ),
1105
+ )
1106
+ parser.add_argument(
1107
+ "--skip-simplify",
1108
+ action="store_true",
1109
+ default=False,
1110
+ help=(
1111
+ "Skip the onnxsim simplification step. "
1112
+ "By default the exported ONNX is simplified in-place with "
1113
+ "onnx-simplifier (pip install onnx-simplifier). "
1114
+ "Use this flag to skip if onnxsim is unavailable or causing issues."
1115
+ ),
1116
+ )
1117
+ parser.add_argument(
1118
+ "--quiet",
1119
+ action="store_true",
1120
+ default=False,
1121
+ help="Suppress verbose progress messages.",
1122
+ )
1123
+
1124
+ # ── Export method ─────────────────────────────────────────────────────────
1125
+ parser.add_argument(
1126
+ "--method",
1127
+ choices=["torch", "optimum"],
1128
+ default="torch",
1129
+ metavar="METHOD",
1130
+ help=(
1131
+ "Export method. "
1132
+ "'torch' (default): clones the official GitHub repo and downloads "
1133
+ "weights from Google Drive via gdown. "
1134
+ "'optimum': downloads from HuggingFace Hub using the transformers "
1135
+ "library β€” proxy-friendly, no CUDA ops, no Google Drive required."
1136
+ ),
1137
+ )
1138
+
1139
+ # ── Utility ───────────────────────────────────────────────────────────────
1140
+ parser.add_argument(
1141
+ "--list-models",
1142
+ action="store_true",
1143
+ default=False,
1144
+ help="Print the model catalogue table and exit.",
1145
+ )
1146
+
1147
+ return parser
1148
+
1149
+
1150
+ # ─────────────────────────────────────────────
1151
+ # Entry point
1152
+ # ─────────────────────────────────────────────
1153
+
1154
+ def main() -> None:
1155
+ parser = build_parser()
1156
+ args = parser.parse_args()
1157
+
1158
+ if args.list_models:
1159
+ print_model_table()
1160
+ return
1161
+
1162
+ if "all" in args.model:
1163
+ shape_tag = f"_{args.shape[0]}x{args.shape[1]}" if args.shape else ""
1164
+ output_dir = os.path.abspath(args.output_dir)
1165
+ pending = [
1166
+ k for k in MODEL_CATALOG
1167
+ if not os.path.exists(os.path.join(output_dir, f"{k}{shape_tag}.onnx"))
1168
+ ]
1169
+ if not pending:
1170
+ print("[INFO] All models already exported. Use --force to re-export.")
1171
+ return
1172
+ skipped = [k for k in MODEL_CATALOG if k not in pending]
1173
+ if skipped:
1174
+ print("[INFO] Already exported (skipping):")
1175
+ for k in skipped:
1176
+ print(f" {k}")
1177
+ print("[INFO] Will export:")
1178
+ for k in pending:
1179
+ print(f" {k}")
1180
+ print()
1181
+ args.model = pending
1182
+
1183
+ if args.weights and len(args.model) > 1:
1184
+ print(
1185
+ "[WARN] --weights applies the same checkpoint to every model in "
1186
+ "--model.\n This is unusual; pass a single --model variant "
1187
+ "when using custom weights."
1188
+ )
1189
+
1190
+ if args.weights and args.method == "optimum":
1191
+ print("[WARN] --weights is ignored with --method optimum. "
1192
+ "HuggingFace weights are always downloaded from the Hub.\n")
1193
+
1194
+ # Install dependencies appropriate to the chosen method
1195
+ if args.method == "optimum":
1196
+ ensure_hf_dependencies()
1197
+ else:
1198
+ ensure_dependencies()
1199
+
1200
+ shape = (args.shape[0], args.shape[1]) if args.shape else None
1201
+ output_dir = os.path.abspath(args.output_dir)
1202
+
1203
+ exported: list[str] = []
1204
+ failed: list[str] = []
1205
+
1206
+ for model_key in args.model:
1207
+ # if MODEL_CATALOG[model_key]["dilation"]:
1208
+ # print(
1209
+ # f"\n[WARN] '{model_key}' is a DC5 (dilation) variant and is "
1210
+ # "temporarily disabled because TIDL does not support dilated "
1211
+ # "convolution in ResNet. Skipping.\n"
1212
+ # )
1213
+ # continue
1214
+
1215
+ print(f"\n{'='*60}")
1216
+ print(f" Exporting: {model_key} [method={args.method}]")
1217
+ print(f"{'='*60}\n")
1218
+
1219
+ try:
1220
+ if args.method == "optimum":
1221
+ out_path = export_model_optimum(
1222
+ model_key = model_key,
1223
+ output_dir = output_dir,
1224
+ shape = shape,
1225
+ opset = args.opset,
1226
+ batch_size = args.batch_size,
1227
+ force = args.force,
1228
+ skip_simplify = args.skip_simplify,
1229
+ )
1230
+ else:
1231
+ out_path = export_model(
1232
+ model_key = model_key,
1233
+ output_dir = output_dir,
1234
+ shape = shape,
1235
+ opset = args.opset,
1236
+ batch_size = args.batch_size,
1237
+ verbose = not args.quiet,
1238
+ custom_weights= args.weights,
1239
+ force = args.force,
1240
+ force_reclone = args.force_reclone,
1241
+ skip_simplify = args.skip_simplify,
1242
+ )
1243
+ exported.append(out_path)
1244
+ except SystemExit:
1245
+ raise
1246
+ except Exception as exc:
1247
+ print(f"[ERROR] Export failed for '{model_key}': {exc}")
1248
+ import traceback
1249
+ traceback.print_exc()
1250
+ failed.append(model_key)
1251
+
1252
+ print("\n" + "=" * 60)
1253
+ print(" Export Summary")
1254
+ print("=" * 60)
1255
+ for path in exported:
1256
+ size_mb = os.path.getsize(path) / (1024 * 1024)
1257
+ print(f" βœ” {os.path.basename(path)} ({size_mb:.1f} MB)")
1258
+ print(f" {path}")
1259
+ if failed:
1260
+ for key in failed:
1261
+ print(f" ✘ {key} (FAILED)")
1262
+ print("=" * 60 + "\n")
1263
+
1264
+ if failed and args.method == "torch":
1265
+ failed_str = " ".join(failed)
1266
+ print(
1267
+ "[TIP] The torch method failed (common causes: Google Drive blocked\n"
1268
+ " by a corporate proxy, or missing CUDA ops).\n"
1269
+ " Try the HuggingFace-based export instead β€” it downloads from\n"
1270
+ " HuggingFace Hub and requires no Google Drive access:\n"
1271
+ f"\n"
1272
+ f" python prepare_model.py --method optimum --model {failed_str}\n"
1273
+ )
1274
+ elif failed and args.method == "optimum":
1275
+ failed_str = " ".join(failed)
1276
+ print(
1277
+ "[TIP] The optimum method failed.\n"
1278
+ " If HuggingFace Hub is accessible, check your transformers\n"
1279
+ " installation. You can also try the torch method with a\n"
1280
+ " manually downloaded checkpoint:\n"
1281
+ f"\n"
1282
+ f" python prepare_model.py --method torch --model {failed_str}\n"
1283
+ f" python prepare_model.py --method torch --model {failed_str} "
1284
+ f"--weights /path/to/checkpoint.pth\n"
1285
+ )
1286
+
1287
+ if failed:
1288
+ sys.exit(1)
1289
+
1290
+
1291
+ if __name__ == "__main__":
1292
+ main()