nitid 0.9.1__py3-none-any.whl
This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
- dfine/__init__.py +23 -0
- dfine/data/__init__.py +8 -0
- dfine/data/coco_names.yml +84 -0
- dfine/exporter.py +399 -0
- dfine/integrations/__init__.py +6 -0
- dfine/integrations/mlflow.py +198 -0
- dfine/integrations/wandb.py +152 -0
- dfine/media.py +130 -0
- dfine/model.py +834 -0
- dfine/nitid.py +153 -0
- dfine/nn/__init__.py +0 -0
- dfine/nn/architecture/__init__.py +15 -0
- dfine/nn/architecture/backbone.py +578 -0
- dfine/nn/architecture/common.py +122 -0
- dfine/nn/architecture/decoder.py +1285 -0
- dfine/nn/architecture/encoder.py +494 -0
- dfine/nn/architecture/model.py +51 -0
- dfine/nn/architecture/ops.py +478 -0
- dfine/nn/build.py +56 -0
- dfine/nn/configs.py +330 -0
- dfine/nn/criterion.py +15 -0
- dfine/nn/distributed.py +30 -0
- dfine/nn/losses/__init__.py +7 -0
- dfine/nn/losses/criterion.py +904 -0
- dfine/nn/losses/matcher.py +264 -0
- dfine/nn/losses/semantic.py +124 -0
- dfine/nn/native_build.py +236 -0
- dfine/nn/openvino_runtime.py +133 -0
- dfine/nn/postprocessor.py +165 -0
- dfine/nn/transfer.py +132 -0
- dfine/plotting.py +66 -0
- dfine/predictor.py +475 -0
- dfine/results.py +535 -0
- dfine/tasks.py +66 -0
- dfine/tracking.py +364 -0
- dfine/trainer.py +1894 -0
- dfine/training_recipes.py +311 -0
- dfine/utils/__init__.py +0 -0
- dfine/utils/augmentations.py +488 -0
- dfine/utils/checkpoint.py +138 -0
- dfine/utils/data.py +1124 -0
- dfine/utils/dataset_converter.py +221 -0
- dfine/utils/device.py +25 -0
- dfine/utils/downloads.py +324 -0
- dfine/utils/logging.py +19 -0
- dfine/utils/ops.py +83 -0
- dfine/utils/reporting.py +114 -0
- dfine/utils/runs.py +151 -0
- dfine/utils/sources.py +383 -0
- dfine/validator.py +904 -0
- nitid/__init__.py +15 -0
- nitid/cli.py +582 -0
- nitid/convert_checkpoint.py +162 -0
- nitid-0.9.1.dist-info/METADATA +446 -0
- nitid-0.9.1.dist-info/RECORD +60 -0
- nitid-0.9.1.dist-info/WHEEL +5 -0
- nitid-0.9.1.dist-info/entry_points.txt +3 -0
- nitid-0.9.1.dist-info/licenses/LICENSE +190 -0
- nitid-0.9.1.dist-info/licenses/THIRD_PARTY_NOTICES.md +95 -0
- nitid-0.9.1.dist-info/top_level.txt +2 -0
dfine/__init__.py
ADDED
|
@@ -0,0 +1,23 @@
|
|
|
1
|
+
"""Compatibility namespace for existing ``dfine`` imports.
|
|
2
|
+
|
|
3
|
+
New applications should import :class:`nitid.NITID` from :mod:`nitid`.
|
|
4
|
+
"""
|
|
5
|
+
|
|
6
|
+
from dfine.media import Frame, FrameMetadata, FrameSink, FrameSource
|
|
7
|
+
from dfine.model import DFINE
|
|
8
|
+
from dfine.nitid import NITID
|
|
9
|
+
from dfine.results import SemanticMask
|
|
10
|
+
from dfine.utils.reporting import BugReport, bugreport
|
|
11
|
+
|
|
12
|
+
__version__ = "0.9.1"
|
|
13
|
+
__all__ = [
|
|
14
|
+
"BugReport",
|
|
15
|
+
"DFINE",
|
|
16
|
+
"NITID",
|
|
17
|
+
"Frame",
|
|
18
|
+
"FrameMetadata",
|
|
19
|
+
"FrameSink",
|
|
20
|
+
"FrameSource",
|
|
21
|
+
"SemanticMask",
|
|
22
|
+
"bugreport",
|
|
23
|
+
]
|
dfine/data/__init__.py
ADDED
|
@@ -0,0 +1,84 @@
|
|
|
1
|
+
# COCO class names embedded in the official pretrained checkpoints.
|
|
2
|
+
# Keep in sync with configs/datasets/coco.yml (checked by tests/unit/test_downloads.py).
|
|
3
|
+
nc: 80
|
|
4
|
+
names:
|
|
5
|
+
0: person
|
|
6
|
+
1: bicycle
|
|
7
|
+
2: car
|
|
8
|
+
3: motorcycle
|
|
9
|
+
4: airplane
|
|
10
|
+
5: bus
|
|
11
|
+
6: train
|
|
12
|
+
7: truck
|
|
13
|
+
8: boat
|
|
14
|
+
9: traffic light
|
|
15
|
+
10: fire hydrant
|
|
16
|
+
11: stop sign
|
|
17
|
+
12: parking meter
|
|
18
|
+
13: bench
|
|
19
|
+
14: bird
|
|
20
|
+
15: cat
|
|
21
|
+
16: dog
|
|
22
|
+
17: horse
|
|
23
|
+
18: sheep
|
|
24
|
+
19: cow
|
|
25
|
+
20: elephant
|
|
26
|
+
21: bear
|
|
27
|
+
22: zebra
|
|
28
|
+
23: giraffe
|
|
29
|
+
24: backpack
|
|
30
|
+
25: umbrella
|
|
31
|
+
26: handbag
|
|
32
|
+
27: tie
|
|
33
|
+
28: suitcase
|
|
34
|
+
29: frisbee
|
|
35
|
+
30: skis
|
|
36
|
+
31: snowboard
|
|
37
|
+
32: sports ball
|
|
38
|
+
33: kite
|
|
39
|
+
34: baseball bat
|
|
40
|
+
35: baseball glove
|
|
41
|
+
36: skateboard
|
|
42
|
+
37: surfboard
|
|
43
|
+
38: tennis racket
|
|
44
|
+
39: bottle
|
|
45
|
+
40: wine glass
|
|
46
|
+
41: cup
|
|
47
|
+
42: fork
|
|
48
|
+
43: knife
|
|
49
|
+
44: spoon
|
|
50
|
+
45: bowl
|
|
51
|
+
46: banana
|
|
52
|
+
47: apple
|
|
53
|
+
48: sandwich
|
|
54
|
+
49: orange
|
|
55
|
+
50: broccoli
|
|
56
|
+
51: carrot
|
|
57
|
+
52: hot dog
|
|
58
|
+
53: pizza
|
|
59
|
+
54: donut
|
|
60
|
+
55: cake
|
|
61
|
+
56: chair
|
|
62
|
+
57: couch
|
|
63
|
+
58: potted plant
|
|
64
|
+
59: bed
|
|
65
|
+
60: dining table
|
|
66
|
+
61: toilet
|
|
67
|
+
62: tv
|
|
68
|
+
63: laptop
|
|
69
|
+
64: mouse
|
|
70
|
+
65: remote
|
|
71
|
+
66: keyboard
|
|
72
|
+
67: cell phone
|
|
73
|
+
68: microwave
|
|
74
|
+
69: oven
|
|
75
|
+
70: toaster
|
|
76
|
+
71: sink
|
|
77
|
+
72: refrigerator
|
|
78
|
+
73: book
|
|
79
|
+
74: clock
|
|
80
|
+
75: vase
|
|
81
|
+
76: scissors
|
|
82
|
+
77: teddy bear
|
|
83
|
+
78: hair drier
|
|
84
|
+
79: toothbrush
|
dfine/exporter.py
ADDED
|
@@ -0,0 +1,399 @@
|
|
|
1
|
+
"""
|
|
2
|
+
DFINEExporter — model export to ONNX, OpenVINO, TensorRT, TorchScript.
|
|
3
|
+
Called internally by DFINE.export(). Not part of the public API.
|
|
4
|
+
"""
|
|
5
|
+
|
|
6
|
+
from __future__ import annotations
|
|
7
|
+
|
|
8
|
+
import os
|
|
9
|
+
import tempfile
|
|
10
|
+
from pathlib import Path
|
|
11
|
+
|
|
12
|
+
import torch
|
|
13
|
+
|
|
14
|
+
from dfine.tasks import normalize_task
|
|
15
|
+
from dfine.utils.logging import LOGGER
|
|
16
|
+
from dfine.utils.runs import atomic_output_path, resolve_run_dir, write_run_metadata
|
|
17
|
+
|
|
18
|
+
# Raw (undecoded) model output names per task — the dict keys the deployed
|
|
19
|
+
# model itself returns, as opposed to the postprocessed (labels, boxes,
|
|
20
|
+
# scores) contract used by the public export() formats.
|
|
21
|
+
_RAW_OUTPUT_NAMES: dict[str, list[str]] = {
|
|
22
|
+
"detect": ["pred_logits", "pred_boxes"],
|
|
23
|
+
"segment": ["pred_logits", "pred_boxes", "pred_masks"],
|
|
24
|
+
"semantic": ["sem_seg_logits"],
|
|
25
|
+
}
|
|
26
|
+
|
|
27
|
+
|
|
28
|
+
# Postprocessed output names per task — the decoded contract produced by
|
|
29
|
+
# DeployModel (model + postprocessor in deploy mode).
|
|
30
|
+
_POSTPROCESSED_OUTPUT_NAMES: dict[str, list[str]] = {
|
|
31
|
+
"detect": ["labels", "boxes", "scores"],
|
|
32
|
+
"segment": ["labels", "boxes", "scores", "masks"],
|
|
33
|
+
"semantic": ["semantic_logits"],
|
|
34
|
+
}
|
|
35
|
+
|
|
36
|
+
|
|
37
|
+
class DeployModel(torch.nn.Module):
|
|
38
|
+
def __init__(self, model, postprocessor, *, semantic: bool = False) -> None:
|
|
39
|
+
super().__init__()
|
|
40
|
+
self.model = model
|
|
41
|
+
self.postprocessor = postprocessor
|
|
42
|
+
self.semantic = semantic
|
|
43
|
+
self.eval()
|
|
44
|
+
|
|
45
|
+
def forward(self, images):
|
|
46
|
+
outputs = self.model(images)
|
|
47
|
+
if self.semantic:
|
|
48
|
+
return outputs["sem_seg_logits"]
|
|
49
|
+
B = images.shape[0]
|
|
50
|
+
H = images.shape[2]
|
|
51
|
+
W = images.shape[3]
|
|
52
|
+
h_t = torch.as_tensor(H, dtype=torch.float32, device=images.device)
|
|
53
|
+
w_t = torch.as_tensor(W, dtype=torch.float32, device=images.device)
|
|
54
|
+
orig_target_sizes = torch.stack([w_t, h_t]).unsqueeze(0).repeat(B, 1)
|
|
55
|
+
return self.postprocessor(outputs, orig_target_sizes)
|
|
56
|
+
|
|
57
|
+
|
|
58
|
+
class DFINEExporter:
|
|
59
|
+
def __init__(self, model, cfg: dict, device: str) -> None:
|
|
60
|
+
self.model = model
|
|
61
|
+
self.cfg = cfg
|
|
62
|
+
self.device = device
|
|
63
|
+
|
|
64
|
+
def export(
|
|
65
|
+
self,
|
|
66
|
+
format: str,
|
|
67
|
+
imgsz: int,
|
|
68
|
+
batch: int,
|
|
69
|
+
dynamic: bool,
|
|
70
|
+
simplify: bool,
|
|
71
|
+
opset: int,
|
|
72
|
+
half: bool,
|
|
73
|
+
postprocess: bool,
|
|
74
|
+
verbose: bool,
|
|
75
|
+
project: str,
|
|
76
|
+
name: str,
|
|
77
|
+
save_dir: str | Path | None,
|
|
78
|
+
output: str | Path | None,
|
|
79
|
+
exist_ok: bool,
|
|
80
|
+
) -> Path:
|
|
81
|
+
format = format.lower()
|
|
82
|
+
suffixes = {
|
|
83
|
+
"onnx": ".onnx",
|
|
84
|
+
"openvino": ".xml",
|
|
85
|
+
"tensorrt": ".engine",
|
|
86
|
+
"torchscript": ".torchscript",
|
|
87
|
+
}
|
|
88
|
+
if format not in suffixes:
|
|
89
|
+
raise ValueError(
|
|
90
|
+
"Unsupported export format: "
|
|
91
|
+
f"{format!r}. Choose: onnx, openvino, tensorrt, torchscript"
|
|
92
|
+
)
|
|
93
|
+
if output is not None:
|
|
94
|
+
out = Path(output)
|
|
95
|
+
if out.exists() and not exist_ok:
|
|
96
|
+
raise FileExistsError(f"Export output already exists: '{out}'")
|
|
97
|
+
run_dir = out.parent
|
|
98
|
+
run_dir.mkdir(parents=True, exist_ok=True)
|
|
99
|
+
else:
|
|
100
|
+
run_dir = resolve_run_dir(
|
|
101
|
+
project=project, name=name, save_dir=save_dir, exist_ok=exist_ok
|
|
102
|
+
)
|
|
103
|
+
out = run_dir / f"dfine_{imgsz}{suffixes[format]}"
|
|
104
|
+
if format == "openvino":
|
|
105
|
+
if out.suffix.lower() != ".xml":
|
|
106
|
+
raise ValueError("OpenVINO output must use the .xml suffix")
|
|
107
|
+
weights_out = out.with_suffix(".bin")
|
|
108
|
+
if weights_out.exists() and not exist_ok:
|
|
109
|
+
raise FileExistsError(f"Export output already exists: '{weights_out}'")
|
|
110
|
+
write_run_metadata(
|
|
111
|
+
run_dir,
|
|
112
|
+
{
|
|
113
|
+
"mode": "export",
|
|
114
|
+
"task": str(self.cfg.get("task", "detect")),
|
|
115
|
+
"format": format,
|
|
116
|
+
"imgsz": imgsz,
|
|
117
|
+
"batch": batch,
|
|
118
|
+
"dynamic": dynamic,
|
|
119
|
+
"simplify": simplify,
|
|
120
|
+
"opset": opset,
|
|
121
|
+
"half": half,
|
|
122
|
+
"postprocess": postprocess,
|
|
123
|
+
"project": project,
|
|
124
|
+
"name": name,
|
|
125
|
+
"save_dir": str(run_dir),
|
|
126
|
+
"output": str(out),
|
|
127
|
+
"exist_ok": exist_ok,
|
|
128
|
+
"verbose": verbose,
|
|
129
|
+
},
|
|
130
|
+
)
|
|
131
|
+
if format == "onnx":
|
|
132
|
+
return self._to_onnx(
|
|
133
|
+
out, imgsz, batch, dynamic, simplify, opset, verbose, postprocess=postprocess
|
|
134
|
+
)
|
|
135
|
+
if format == "openvino":
|
|
136
|
+
return self._to_openvino(
|
|
137
|
+
out,
|
|
138
|
+
imgsz,
|
|
139
|
+
batch,
|
|
140
|
+
dynamic,
|
|
141
|
+
simplify,
|
|
142
|
+
opset,
|
|
143
|
+
half,
|
|
144
|
+
verbose,
|
|
145
|
+
postprocess=postprocess,
|
|
146
|
+
)
|
|
147
|
+
if format == "tensorrt":
|
|
148
|
+
return self._to_tensorrt(
|
|
149
|
+
out, imgsz, batch, dynamic, half, verbose, postprocess=postprocess
|
|
150
|
+
)
|
|
151
|
+
if format == "torchscript":
|
|
152
|
+
return self._to_torchscript(out, imgsz, batch, verbose)
|
|
153
|
+
raise AssertionError("unreachable")
|
|
154
|
+
|
|
155
|
+
# ── ONNX ────────────────────────────────────────────────────────────────
|
|
156
|
+
|
|
157
|
+
def _to_onnx(
|
|
158
|
+
self, out, imgsz, batch, dynamic, simplify, opset, verbose, *, postprocess: bool = True
|
|
159
|
+
) -> Path:
|
|
160
|
+
import onnx
|
|
161
|
+
|
|
162
|
+
with atomic_output_path(out) as temporary:
|
|
163
|
+
self._export_onnx_to_path(
|
|
164
|
+
temporary, imgsz, batch, dynamic, opset, postprocess=postprocess
|
|
165
|
+
)
|
|
166
|
+
if simplify:
|
|
167
|
+
import onnxsim
|
|
168
|
+
|
|
169
|
+
model_onnx = onnx.load(str(temporary))
|
|
170
|
+
model_onnx, ok = onnxsim.simplify(model_onnx)
|
|
171
|
+
if ok:
|
|
172
|
+
onnx.save(model_onnx, str(temporary))
|
|
173
|
+
if verbose:
|
|
174
|
+
LOGGER.info("ONNX export saved to %s", out)
|
|
175
|
+
return out
|
|
176
|
+
|
|
177
|
+
# ── OpenVINO ────────────────────────────────────────────────────────────
|
|
178
|
+
|
|
179
|
+
def _to_openvino(
|
|
180
|
+
self,
|
|
181
|
+
out: Path,
|
|
182
|
+
imgsz: int,
|
|
183
|
+
batch: int,
|
|
184
|
+
dynamic: bool,
|
|
185
|
+
simplify: bool,
|
|
186
|
+
opset: int,
|
|
187
|
+
half: bool,
|
|
188
|
+
verbose: bool,
|
|
189
|
+
*,
|
|
190
|
+
postprocess: bool = True,
|
|
191
|
+
) -> Path:
|
|
192
|
+
"""Export the corrected ONNX graph, then convert it to OpenVINO IR."""
|
|
193
|
+
try:
|
|
194
|
+
import openvino as ov
|
|
195
|
+
except ImportError:
|
|
196
|
+
raise ImportError(
|
|
197
|
+
"OpenVINO is not installed. Install it with: uv sync --extra openvino"
|
|
198
|
+
) from None
|
|
199
|
+
|
|
200
|
+
out.parent.mkdir(parents=True, exist_ok=True)
|
|
201
|
+
with tempfile.TemporaryDirectory(prefix=f".{out.stem}.", dir=out.parent) as directory:
|
|
202
|
+
staging_dir = Path(directory)
|
|
203
|
+
onnx_path = staging_dir / "model.onnx"
|
|
204
|
+
ir_path = staging_dir / "model.xml"
|
|
205
|
+
self._to_onnx(
|
|
206
|
+
onnx_path,
|
|
207
|
+
imgsz,
|
|
208
|
+
batch,
|
|
209
|
+
dynamic,
|
|
210
|
+
simplify,
|
|
211
|
+
opset,
|
|
212
|
+
verbose=False,
|
|
213
|
+
postprocess=postprocess,
|
|
214
|
+
)
|
|
215
|
+
|
|
216
|
+
ov_model = ov.convert_model(onnx_path)
|
|
217
|
+
ov.save_model(ov_model, ir_path, compress_to_fp16=half)
|
|
218
|
+
|
|
219
|
+
# Ensure the serialized IR can be loaded and compiled before publishing it.
|
|
220
|
+
core = ov.Core()
|
|
221
|
+
core.compile_model(core.read_model(ir_path), "CPU")
|
|
222
|
+
|
|
223
|
+
staged_weights = ir_path.with_suffix(".bin")
|
|
224
|
+
if not staged_weights.is_file():
|
|
225
|
+
raise RuntimeError("OpenVINO conversion did not produce the expected .bin file")
|
|
226
|
+
for staged_file in (staged_weights, ir_path):
|
|
227
|
+
with staged_file.open("rb") as stream:
|
|
228
|
+
os.fsync(stream.fileno())
|
|
229
|
+
|
|
230
|
+
# Publish weights first and XML last, so a visible XML always has complete weights.
|
|
231
|
+
os.replace(staged_weights, out.with_suffix(".bin"))
|
|
232
|
+
os.replace(ir_path, out)
|
|
233
|
+
|
|
234
|
+
LOGGER.info("OpenVINO IR export saved to %s", out)
|
|
235
|
+
return out
|
|
236
|
+
|
|
237
|
+
def _export_onnx_to_path(
|
|
238
|
+
self,
|
|
239
|
+
path: Path,
|
|
240
|
+
imgsz: int,
|
|
241
|
+
batch: int,
|
|
242
|
+
dynamic: bool,
|
|
243
|
+
opset: int,
|
|
244
|
+
*,
|
|
245
|
+
postprocess: bool = True,
|
|
246
|
+
) -> None:
|
|
247
|
+
"""
|
|
248
|
+
Trace the model to ONNX at an explicit output path.
|
|
249
|
+
|
|
250
|
+
With ``postprocess=False`` the deployed model is traced on its own, so the
|
|
251
|
+
graph exposes the decoder's raw outputs (``pred_logits``/``pred_boxes``/…)
|
|
252
|
+
instead of the decoded ``(labels, boxes, scores, …)`` contract.
|
|
253
|
+
"""
|
|
254
|
+
from typing import Any
|
|
255
|
+
|
|
256
|
+
from dfine.nn.build import build_postprocessor
|
|
257
|
+
|
|
258
|
+
task = normalize_task(str(self.cfg.get("task", "detect")))
|
|
259
|
+
|
|
260
|
+
if postprocess:
|
|
261
|
+
postprocessor: Any = build_postprocessor(self.cfg)
|
|
262
|
+
if hasattr(postprocessor, "deploy"):
|
|
263
|
+
postprocessor.deploy()
|
|
264
|
+
postprocessor.to(self.device)
|
|
265
|
+
traced_model: torch.nn.Module = DeployModel(
|
|
266
|
+
self.model, postprocessor, semantic=task == "semantic"
|
|
267
|
+
)
|
|
268
|
+
traced_model.eval()
|
|
269
|
+
output_names = _POSTPROCESSED_OUTPUT_NAMES[task]
|
|
270
|
+
else:
|
|
271
|
+
traced_model, output_names = self.prepare_raw_trace(imgsz)
|
|
272
|
+
|
|
273
|
+
dummy = torch.zeros(batch, 3, imgsz, imgsz, device=self.device)
|
|
274
|
+
dynamic_axes = (
|
|
275
|
+
{name: {0: "batch"} for name in ("images", *output_names)} if dynamic else None
|
|
276
|
+
)
|
|
277
|
+
|
|
278
|
+
torch.onnx.export(
|
|
279
|
+
traced_model,
|
|
280
|
+
(dummy,),
|
|
281
|
+
str(path),
|
|
282
|
+
dynamo=False,
|
|
283
|
+
opset_version=opset,
|
|
284
|
+
input_names=["images"],
|
|
285
|
+
output_names=output_names,
|
|
286
|
+
dynamic_axes=dynamic_axes,
|
|
287
|
+
)
|
|
288
|
+
|
|
289
|
+
# ── Raw (non-postprocessed) trace — used by the OpenVINO runtime ────────
|
|
290
|
+
|
|
291
|
+
def prepare_raw_trace(self, imgsz: int) -> tuple[torch.nn.Module, list[str]]:
|
|
292
|
+
"""
|
|
293
|
+
Return ``(model, output_names)`` for a raw trace of the deployed model
|
|
294
|
+
at ``imgsz`` — its own ``pred_logits``/``pred_boxes``/etc. dict, not
|
|
295
|
+
the postprocessed ``(labels, boxes, scores)`` graph ``export()``
|
|
296
|
+
produces by default. Used by ``export(..., postprocess=False)`` and by
|
|
297
|
+
``dfine.nn.openvino_runtime.compile_raw_openvino``.
|
|
298
|
+
"""
|
|
299
|
+
task = normalize_task(str(self.cfg.get("task", "detect")))
|
|
300
|
+
return self.model, _RAW_OUTPUT_NAMES[task]
|
|
301
|
+
|
|
302
|
+
# ── TensorRT ─────────────────────────────────────────────────────────────
|
|
303
|
+
|
|
304
|
+
def _to_tensorrt(
|
|
305
|
+
self,
|
|
306
|
+
out: Path,
|
|
307
|
+
imgsz: int,
|
|
308
|
+
batch: int,
|
|
309
|
+
dynamic: bool,
|
|
310
|
+
half: bool,
|
|
311
|
+
verbose: bool,
|
|
312
|
+
*,
|
|
313
|
+
postprocess: bool = True,
|
|
314
|
+
) -> Path:
|
|
315
|
+
"""
|
|
316
|
+
Export to a TensorRT serialised engine via the TRT Python API.
|
|
317
|
+
|
|
318
|
+
Workflow: trace model → temp ONNX → parse with TRT → write .engine file.
|
|
319
|
+
Requires ``tensorrt`` (``uv sync --extra tensorrt``).
|
|
320
|
+
"""
|
|
321
|
+
import tempfile
|
|
322
|
+
|
|
323
|
+
try:
|
|
324
|
+
import tensorrt as trt
|
|
325
|
+
except ImportError:
|
|
326
|
+
raise ImportError(
|
|
327
|
+
"TensorRT is not installed. Install it with: uv sync --extra tensorrt"
|
|
328
|
+
) from None
|
|
329
|
+
|
|
330
|
+
# ── Step 1: trace to a temporary ONNX ────────────────────────────────
|
|
331
|
+
with tempfile.NamedTemporaryFile(suffix=".onnx", delete=False) as f:
|
|
332
|
+
tmp_onnx = Path(f.name)
|
|
333
|
+
|
|
334
|
+
try:
|
|
335
|
+
self._export_onnx_to_path(
|
|
336
|
+
tmp_onnx, imgsz, batch, dynamic, opset=17, postprocess=postprocess
|
|
337
|
+
)
|
|
338
|
+
|
|
339
|
+
# ── Step 2: build TRT engine from ONNX ───────────────────────────
|
|
340
|
+
trt_logger = trt.Logger(trt.Logger.INFO if verbose else trt.Logger.WARNING)
|
|
341
|
+
builder = trt.Builder(trt_logger)
|
|
342
|
+
network = builder.create_network(
|
|
343
|
+
1 << int(trt.NetworkDefinitionCreationFlag.EXPLICIT_BATCH)
|
|
344
|
+
)
|
|
345
|
+
parser = trt.OnnxParser(network, trt_logger)
|
|
346
|
+
|
|
347
|
+
with open(str(tmp_onnx), "rb") as f:
|
|
348
|
+
if not parser.parse(f.read()):
|
|
349
|
+
errors = "\n".join(str(parser.get_error(i)) for i in range(parser.num_errors))
|
|
350
|
+
raise RuntimeError(f"TensorRT failed to parse ONNX:\n{errors}")
|
|
351
|
+
|
|
352
|
+
config = builder.create_builder_config()
|
|
353
|
+
|
|
354
|
+
# Workspace: 1 GB — handle API change between TRT 8.x and 8.5+
|
|
355
|
+
workspace = 1 << 30
|
|
356
|
+
if hasattr(trt, "MemoryPoolType"):
|
|
357
|
+
config.set_memory_pool_limit(trt.MemoryPoolType.WORKSPACE, workspace)
|
|
358
|
+
else:
|
|
359
|
+
config.max_workspace_size = workspace # type: ignore[attr-defined]
|
|
360
|
+
|
|
361
|
+
if half:
|
|
362
|
+
if not builder.platform_has_fast_fp16:
|
|
363
|
+
LOGGER.warning("half=True but GPU reports no native FP16 support")
|
|
364
|
+
config.set_flag(trt.BuilderFlag.FP16)
|
|
365
|
+
|
|
366
|
+
if dynamic:
|
|
367
|
+
profile = builder.create_optimization_profile()
|
|
368
|
+
# min=1, opt=batch, max=batch*4
|
|
369
|
+
profile.set_shape(
|
|
370
|
+
"images",
|
|
371
|
+
(1, 3, imgsz, imgsz),
|
|
372
|
+
(batch, 3, imgsz, imgsz),
|
|
373
|
+
(batch * 4, 3, imgsz, imgsz),
|
|
374
|
+
)
|
|
375
|
+
config.add_optimization_profile(profile)
|
|
376
|
+
|
|
377
|
+
engine_bytes = builder.build_serialized_network(network, config)
|
|
378
|
+
if engine_bytes is None:
|
|
379
|
+
raise RuntimeError("TensorRT engine build failed — check GPU and TRT logs")
|
|
380
|
+
|
|
381
|
+
with atomic_output_path(out) as temporary:
|
|
382
|
+
temporary.write_bytes(engine_bytes)
|
|
383
|
+
|
|
384
|
+
finally:
|
|
385
|
+
tmp_onnx.unlink(missing_ok=True)
|
|
386
|
+
|
|
387
|
+
LOGGER.info(f"TensorRT engine saved to {out}")
|
|
388
|
+
return out
|
|
389
|
+
|
|
390
|
+
# ── TorchScript ──────────────────────────────────────────────────────────
|
|
391
|
+
|
|
392
|
+
def _to_torchscript(self, out, imgsz, batch, verbose) -> Path:
|
|
393
|
+
dummy = torch.zeros(batch, 3, imgsz, imgsz, device=self.device)
|
|
394
|
+
# D-FINE forward returns a dict; strict=False allows tracing dict outputs
|
|
395
|
+
scripted = torch.jit.trace(self.model, dummy, strict=False)
|
|
396
|
+
with atomic_output_path(out) as temporary:
|
|
397
|
+
scripted.save(str(temporary))
|
|
398
|
+
LOGGER.info(f"TorchScript export saved to {out}")
|
|
399
|
+
return out
|