easydetect 0.1.0__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.
easydetect/__init__.py ADDED
@@ -0,0 +1,44 @@
1
+ # Apache-2.0
2
+ """easydetect — real-time object detection you can ship, 100% Apache-2.0.
3
+
4
+ pip install easydetect
5
+
6
+ from easydetect import Detector
7
+
8
+ model = Detector("dfine-s") # weights fetched from the mirror
9
+ results = model("photo.jpg") # list[Results]
10
+ results[0].boxes.xyxy, results[0].boxes.conf # plain numpy
11
+ results[0].save()
12
+
13
+ model.train(data="data.yaml", epochs=50) # your own dataset
14
+ model.val(data="data.yaml").box.map50 # COCO-style mAP
15
+ model.export(format="openvino", half=True) # IR + labels.txt
16
+
17
+ The detector is D-FINE (Peterande/D-FINE, Apache-2.0), a real-time DETR, in
18
+ five sizes n/s/m/l/x; its released COCO weights load unchanged. The API,
19
+ trainer, validator, exporter, predictor and CLI are this project's own. No AGPL
20
+ code or weights anywhere, so this package can ship inside a product. Inference
21
+ needs numpy/opencv/openvino/pyyaml; training adds
22
+ ``pip install "easydetect[train]"`` (torch, torchvision, scipy, onnx).
23
+ """
24
+
25
+ from __future__ import annotations
26
+
27
+ __version__ = "0.1.0"
28
+
29
+ from .errors import DownloadError, EasyDetectError, ModelNotFoundError
30
+ from .metrics import BoxMetrics, DetMetrics
31
+ from .model import Detector
32
+ from .results import Boxes, Results
33
+
34
+ __all__ = [
35
+ "Detector",
36
+ "Results",
37
+ "Boxes",
38
+ "DetMetrics",
39
+ "BoxMetrics",
40
+ "EasyDetectError",
41
+ "ModelNotFoundError",
42
+ "DownloadError",
43
+ "__version__",
44
+ ]
easydetect/cli.py ADDED
@@ -0,0 +1,149 @@
1
+ # Apache-2.0
2
+ """``easydetect <mode> key=value ...`` — the command line for this package.
3
+
4
+ easydetect predict model=dfine-s source=bus.jpg conf=0.5
5
+ easydetect train model=dfine-s data=data.yaml epochs=100 imgsz=640
6
+ easydetect val model=best.pt data=data.yaml
7
+ easydetect export model=best.pt format=openvino half=true
8
+ easydetect track model=best.pt source=video.mp4
9
+ """
10
+
11
+ from __future__ import annotations
12
+
13
+ import sys
14
+ from typing import Any
15
+
16
+ MODES = ("predict", "track", "train", "val", "export")
17
+
18
+ HELP = """easydetect — real-time object detection (D-FINE), Apache-2.0
19
+
20
+ Usage: easydetect <mode> key=value ...
21
+
22
+ Modes:
23
+ predict run detection on an image / folder / video / url / camera
24
+ track predict and keep an id on each box across frames
25
+ train train on a data.yaml dataset (images/ + labels/*.txt)
26
+ val COCO-style mAP50 / mAP50-95 on the val split
27
+ export write an OpenVINO IR (or ONNX) next to the checkpoint
28
+
29
+ Common keys:
30
+ model=dfine-n|s|m|l|x|best.pt|model.xml source=bus.jpg|dir|video.mp4|0|url
31
+ data=data.yaml epochs=100 imgsz=640 batch=8 conf=0.25 device=0|cpu|AUTO
32
+ project=runs name=predict save=true show=false format=openvino half=true
33
+
34
+ Examples:
35
+ easydetect predict model=dfine-s source=bus.jpg conf=0.5
36
+ easydetect train model=dfine-s data=data.yaml epochs=100
37
+ easydetect val model=best.pt data=data.yaml
38
+ easydetect export model=best.pt format=openvino half=true
39
+ """
40
+
41
+
42
+ def parse_value(text: str) -> Any:
43
+ """``true`` -> True, ``3`` -> 3, ``0.5`` -> 0.5, ``0,1`` -> [0, 1], else str."""
44
+ low = text.lower()
45
+ if low in ("true", "yes"):
46
+ return True
47
+ if low in ("false", "no"):
48
+ return False
49
+ if low in ("none", "null"):
50
+ return None
51
+ for cast in (int, float):
52
+ try:
53
+ return cast(text)
54
+ except ValueError:
55
+ pass
56
+ stripped = text.strip("[]")
57
+ if "," in stripped:
58
+ return [parse_value(part.strip()) for part in stripped.split(",") if part.strip()]
59
+ return text
60
+
61
+
62
+ def parse_args(argv: list[str]) -> tuple[str, dict[str, Any]]:
63
+ """-> ``(mode, overrides)``. The mode may be positional or ``mode=predict``."""
64
+ mode = ""
65
+ overrides: dict[str, Any] = {}
66
+ for arg in argv:
67
+ if "=" in arg:
68
+ key, _, value = arg.partition("=")
69
+ key = key.strip().lstrip("-")
70
+ if key == "mode":
71
+ mode = str(value)
72
+ else:
73
+ overrides[key] = parse_value(value)
74
+ elif not mode:
75
+ mode = arg
76
+ else:
77
+ raise SystemExit(f"unexpected argument {arg!r} — arguments are key=value pairs")
78
+ return mode, overrides
79
+
80
+
81
+ def _take(overrides: dict[str, Any], key: str, default: Any = None) -> Any:
82
+ return overrides.pop(key, default)
83
+
84
+
85
+ def main(argv: list[str] | None = None) -> int:
86
+ argv = list(sys.argv[1:] if argv is None else argv)
87
+ if not argv or argv[0] in ("-h", "--help", "help"):
88
+ print(HELP)
89
+ return 0
90
+ if argv[0] in ("-v", "--version", "version"):
91
+ from . import __version__
92
+
93
+ print(__version__)
94
+ return 0
95
+
96
+ mode, overrides = parse_args(argv)
97
+ if mode not in MODES:
98
+ print(f"unknown mode {mode!r}. Modes: {', '.join(MODES)}\n", file=sys.stderr)
99
+ print(HELP, file=sys.stderr)
100
+ return 2
101
+
102
+ from .model import Detector
103
+
104
+ model_name = _take(overrides, "model", "dfine-s")
105
+ device = _take(overrides, "device")
106
+ verbose = _take(overrides, "verbose", True)
107
+
108
+ if mode in ("predict", "track"):
109
+ model = Detector(model_name, device=device or "AUTO", verbose=verbose)
110
+ source = _take(overrides, "source")
111
+ if source is None:
112
+ print(
113
+ "predict needs source=… (image, folder, video, url, camera index)",
114
+ file=sys.stderr,
115
+ )
116
+ return 2
117
+ overrides.setdefault("save", True)
118
+ overrides["stream"] = False
119
+ runner = model.predict if mode == "predict" else model.track
120
+ results = runner(source, **overrides)
121
+ print(f"{len(results)} result(s)")
122
+ return 0
123
+
124
+ model = Detector(model_name, verbose=verbose)
125
+ if mode == "train":
126
+ data = _take(overrides, "data")
127
+ if data is None:
128
+ print("train needs data=path/to/data.yaml", file=sys.stderr)
129
+ return 2
130
+ best = model.train(data=data, device=device, **overrides)
131
+ print(best)
132
+ return 0
133
+ if mode == "val":
134
+ data = _take(overrides, "data")
135
+ if data is None:
136
+ print("val needs data=path/to/data.yaml", file=sys.stderr)
137
+ return 2
138
+ metrics = model.val(data=data, device=device, **overrides)
139
+ if not verbose: # model.val() already printed the line when verbose
140
+ print(f"mAP50 {metrics.box.map50:.4f} mAP50-95 {metrics.box.map:.4f}")
141
+ return 0
142
+
143
+ path = model.export(**overrides)
144
+ print(path)
145
+ return 0
146
+
147
+
148
+ if __name__ == "__main__": # pragma: no cover
149
+ raise SystemExit(main())
@@ -0,0 +1 @@
1
+ # Apache-2.0
@@ -0,0 +1,161 @@
1
+ # Apache-2.0
2
+ """Detection dataset: a data.yaml plus images/ and labels/*.txt.
3
+
4
+ data.yaml:
5
+ path: dataset root (optional)
6
+ train: images dir or txt list
7
+ val: images dir or txt list
8
+ names: {0: person, ...} or [person, ...]
9
+ Labels: <images-dir with 'images' replaced by 'labels'>/<stem>.txt
10
+ each line: cls cx cy w h (normalized)
11
+ """
12
+
13
+ import random
14
+ from pathlib import Path
15
+
16
+ import cv2
17
+ import numpy as np
18
+ import torch
19
+ import yaml
20
+ from torch.utils.data import Dataset
21
+
22
+ from .labels import label_path, label_row_to_box
23
+
24
+ IMG_EXT = {".jpg", ".jpeg", ".png", ".bmp", ".webp"}
25
+
26
+
27
+ def load_data_yaml(path):
28
+ with open(path, encoding="utf-8") as f:
29
+ cfg = yaml.safe_load(f)
30
+ yaml_dir = Path(path).resolve().parent
31
+ root = Path(cfg.get("path") or yaml_dir)
32
+ if not root.is_absolute():
33
+ root = yaml_dir / root
34
+ if not root.exists():
35
+ # a dataset downloaded from elsewhere often keeps its author's path
36
+ # (/content/datasets/..., a Windows home folder); the yaml's own folder
37
+ # is the only root that can be right on this machine
38
+ root = yaml_dir
39
+ names = cfg["names"]
40
+ if isinstance(names, list):
41
+ names = {i: n for i, n in enumerate(names)}
42
+ names = {int(k): str(v) for k, v in names.items()}
43
+ return {
44
+ "root": root,
45
+ "yaml_dir": yaml_dir,
46
+ "train": cfg.get("train"),
47
+ "val": cfg.get("val"),
48
+ "names": names,
49
+ "nc": len(names),
50
+ }
51
+
52
+
53
+ def _locate(root: Path, yaml_dir: Path, spec: str) -> Path:
54
+ """Where a train/val entry points, trying the ways data.yaml files are written.
55
+
56
+ Relative to ``path:`` first; then to the yaml itself; then with leading
57
+ ``../`` dropped — exports that sit beside their splits still write
58
+ ``train: ../train/images``.
59
+ """
60
+ spec_path = Path(spec)
61
+ if spec_path.is_absolute():
62
+ return spec_path
63
+ tried = [root / spec_path, yaml_dir / spec_path]
64
+ stripped = Path(*[part for part in spec_path.parts if part != ".."] or ["."])
65
+ tried.append(yaml_dir / stripped)
66
+ for candidate in tried:
67
+ if candidate.exists():
68
+ return candidate
69
+ raise FileNotFoundError(
70
+ f"train/val entry {spec!r} not found; tried "
71
+ + ", ".join(str(c) for c in dict.fromkeys(tried))
72
+ )
73
+
74
+
75
+ def _list_images(root: Path, spec, yaml_dir: Path | None = None):
76
+ """Images for one split: a folder, a .txt list, or a list of either."""
77
+ if isinstance(spec, (list, tuple)):
78
+ files = [f for s in spec for f in _list_images(root, s, yaml_dir)]
79
+ return list(dict.fromkeys(files))
80
+ p = _locate(root, yaml_dir or root, str(spec))
81
+ if p.is_dir():
82
+ return sorted(f for f in p.rglob("*") if f.suffix.lower() in IMG_EXT)
83
+ if p.suffix == ".txt":
84
+ base = p.parent
85
+ out = []
86
+ for line in p.read_text().splitlines():
87
+ line = line.strip()
88
+ if line:
89
+ q = Path(line)
90
+ out.append(q if q.is_absolute() else base / q)
91
+ return out
92
+ raise FileNotFoundError(f"train/val entry not found: {p}")
93
+
94
+
95
+ _label_path = label_path # older name, kept for callers
96
+
97
+
98
+ class DetDataset(Dataset):
99
+ def __init__(self, data_yaml, split="train", imgsz=640, augment=True):
100
+ cfg = load_data_yaml(data_yaml)
101
+ self.names, self.nc = cfg["names"], cfg["nc"]
102
+ self.imgsz = imgsz
103
+ self.augment = augment and split == "train"
104
+ if cfg[split] is None:
105
+ raise ValueError(
106
+ f"{data_yaml} has no '{split}:' entry. Add one — it may point at the "
107
+ f"training images, but then mAP only measures memorisation."
108
+ )
109
+ self.files = _list_images(cfg["root"], cfg[split], cfg["yaml_dir"])
110
+ if not self.files:
111
+ raise FileNotFoundError(f"no images for split '{split}'")
112
+
113
+ def __len__(self):
114
+ return len(self.files)
115
+
116
+ def _load_labels(self, img_file):
117
+ lp = _label_path(img_file)
118
+ if not lp.exists():
119
+ return np.zeros((0, 5), np.float32)
120
+ rows = []
121
+ for line in lp.read_text().splitlines():
122
+ row = label_row_to_box(line.split())
123
+ if row is not None:
124
+ rows.append(row)
125
+ return np.asarray(rows, np.float32) if rows else np.zeros((0, 5), np.float32)
126
+
127
+ def __getitem__(self, i):
128
+ f = self.files[i]
129
+ img = cv2.imread(str(f))
130
+ if img is None:
131
+ raise FileNotFoundError(f)
132
+ labels = self._load_labels(f) # cls, cx, cy, w, h (normalized)
133
+
134
+ if self.augment:
135
+ # hflip
136
+ if random.random() < 0.5:
137
+ img = img[:, ::-1]
138
+ if len(labels):
139
+ labels[:, 1] = 1.0 - labels[:, 1]
140
+ # HSV jitter
141
+ hsv = cv2.cvtColor(img, cv2.COLOR_BGR2HSV).astype(np.int16)
142
+ hsv[..., 0] = (hsv[..., 0] + random.randint(-8, 8)) % 180
143
+ hsv[..., 1] = np.clip(hsv[..., 1] + random.randint(-30, 30), 0, 255)
144
+ hsv[..., 2] = np.clip(hsv[..., 2] + random.randint(-30, 30), 0, 255)
145
+ img = cv2.cvtColor(hsv.astype(np.uint8), cv2.COLOR_HSV2BGR)
146
+
147
+ img = cv2.resize(img, (self.imgsz, self.imgsz)) # plain resize, as D-FINE trains
148
+ img = cv2.cvtColor(img, cv2.COLOR_BGR2RGB).astype(np.float32) / 255.0
149
+ tensor = torch.from_numpy(img.transpose(2, 0, 1)).contiguous()
150
+
151
+ target = {
152
+ "labels": torch.as_tensor(labels[:, 0], dtype=torch.long),
153
+ "boxes": torch.as_tensor(labels[:, 1:5], dtype=torch.float32), # cxcywh 0..1
154
+ }
155
+ return tensor, target
156
+
157
+ @staticmethod
158
+ def collate(batch):
159
+ imgs = torch.stack([b[0] for b in batch])
160
+ targets = [b[1] for b in batch]
161
+ return imgs, targets
@@ -0,0 +1,48 @@
1
+ # Apache-2.0
2
+ """Where a label lives, and how one line of it reads — shared by everything.
3
+
4
+ The trainer reads labels, the platform writes them, the exporter packs them.
5
+ If any two of those disagree on where ``a.jpg``'s boxes are, training runs on
6
+ nothing and nobody is told. So there is one rule, here, and no torch import:
7
+ the platform uses it at startup.
8
+ """
9
+
10
+ from __future__ import annotations
11
+
12
+ from pathlib import Path, PurePath
13
+
14
+ import numpy as np
15
+
16
+
17
+ def label_path(image: str | PurePath) -> PurePath:
18
+ """``…/images/train/a.jpg`` -> ``…/labels/train/a.txt``.
19
+
20
+ The last folder called ``images`` becomes ``labels``; with none, the label
21
+ sits beside the image. Works on relative paths too, which is how an export
22
+ lays out a zip that the trainer will read after unpacking.
23
+ """
24
+ image = image if isinstance(image, PurePath) else Path(image)
25
+ parts = list(image.parts)
26
+ for i in range(len(parts) - 2, -1, -1): # folders only, never the file name
27
+ if parts[i] == "images":
28
+ parts[i] = "labels"
29
+ return type(image)(*parts).with_suffix(".txt")
30
+ return image.with_suffix(".txt")
31
+
32
+
33
+ def label_row_to_box(values: list[str]) -> list[float] | None:
34
+ """One label line -> ``[cls, cx, cy, w, h]``, or None for a line to skip.
35
+
36
+ ``cls cx cy w h`` is a box (a trailing sixth value, a confidence, is
37
+ ignored). ``cls x1 y1 x2 y2 x3 y3 ...`` is a segmentation polygon, common
38
+ in exported datasets; reading its first four numbers as a box would train
39
+ on nonsense without a single error, so it becomes its bounding box.
40
+ """
41
+ if len(values) >= 7 and len(values) % 2 == 1:
42
+ xy = np.asarray(values[1:], np.float32).reshape(-1, 2)
43
+ (x0, y0), (x1, y1) = xy.min(0), xy.max(0)
44
+ return [float(values[0]), float(x0 + x1) / 2, float(y0 + y1) / 2,
45
+ float(x1 - x0), float(y1 - y0)]
46
+ if len(values) >= 5:
47
+ return [float(v) for v in values[:5]]
48
+ return None
@@ -0,0 +1,145 @@
1
+ # Apache-2.0
2
+ """Pretrained-weight download and cache — stdlib only (urllib), cached in ~/.easydetect.
3
+
4
+ ``Detector("dfine-s")`` has to just work, so the first predict pulls the
5
+ OpenVINO IR for that name from the public mirror and keeps it under
6
+ ``~/.easydetect/`` (override with ``$EASYDETECT_HOME``). The mirror base URL is
7
+ ``$EASYDETECT_ASSETS_URL``-overridable, which is also how an offline site points
8
+ the package at an internal copy.
9
+ """
10
+
11
+ from __future__ import annotations
12
+
13
+ import os
14
+ import shutil
15
+ import sys
16
+ import urllib.error
17
+ import urllib.request
18
+ from pathlib import Path
19
+
20
+ from .errors import DownloadError, ModelNotFoundError
21
+
22
+ #: Where the released weights live. Layout: ``<base>/<name>/<file>``.
23
+ DEFAULT_ASSETS_URL = "https://huggingface.co/leeyunjai/easydetect/resolve/main"
24
+
25
+ #: Names the mirror knows, with the spellings users type mapped onto them.
26
+ MODEL_NAMES = ("dfine-n", "dfine-s", "dfine-m", "dfine-l", "dfine-x")
27
+
28
+ _TIMEOUT = 30.0
29
+
30
+
31
+ def normalize_name(name: str) -> str:
32
+ """``dfine_s`` / ``D-FINE-S`` / ``DFINE-S`` / ``dfine-s.pt`` -> ``dfine-s``."""
33
+ stem = str(name).strip().lower()
34
+ for suffix in (".pt", ".xml", ".onnx"):
35
+ if stem.endswith(suffix):
36
+ stem = stem[: -len(suffix)]
37
+ stem = stem.replace("_", "-")
38
+ return "dfine-" + stem[len("d-fine-"):] if stem.startswith("d-fine-") else stem
39
+
40
+
41
+ def is_model_name(name: str) -> bool:
42
+ return normalize_name(name) in MODEL_NAMES
43
+
44
+
45
+ def assets_url() -> str:
46
+ return os.environ.get("EASYDETECT_ASSETS_URL", DEFAULT_ASSETS_URL).rstrip("/")
47
+
48
+
49
+ def cache_dir() -> Path:
50
+ root = Path(os.environ.get("EASYDETECT_HOME", Path.home() / ".easydetect")).expanduser()
51
+ root.mkdir(parents=True, exist_ok=True)
52
+ return root
53
+
54
+
55
+ def read_url_bytes(url: str, timeout: float = _TIMEOUT) -> bytes:
56
+ """Fetch a URL into memory (used for image URLs as well as weights)."""
57
+ request = urllib.request.Request(url, headers={"User-Agent": "easydetect"})
58
+ try:
59
+ with urllib.request.urlopen(request, timeout=timeout) as response:
60
+ return response.read()
61
+ except urllib.error.URLError as exc:
62
+ raise DownloadError(f"could not fetch {url}: {exc}") from exc
63
+
64
+
65
+ def download(url: str, dest: Path, progress: bool = True) -> Path:
66
+ """Download ``url`` to ``dest`` (atomically). Existing files are reused."""
67
+ dest = Path(dest)
68
+ if dest.exists() and dest.stat().st_size:
69
+ return dest
70
+ dest.parent.mkdir(parents=True, exist_ok=True)
71
+ tmp = dest.with_suffix(dest.suffix + ".part")
72
+ request = urllib.request.Request(url, headers={"User-Agent": "easydetect"})
73
+ try:
74
+ with urllib.request.urlopen(request, timeout=_TIMEOUT) as response, open(tmp, "wb") as out:
75
+ total = int(response.headers.get("Content-Length") or 0)
76
+ done = 0
77
+ while chunk := response.read(1 << 16):
78
+ out.write(chunk)
79
+ done += len(chunk)
80
+ if progress and total and sys.stderr.isatty():
81
+ pct = 100 * done / total
82
+ print(f"\rdownloading {dest.name} {pct:5.1f}%", end="", file=sys.stderr)
83
+ if progress and sys.stderr.isatty():
84
+ print(f"\rdownloading {dest.name} done ", file=sys.stderr)
85
+ except urllib.error.HTTPError as exc:
86
+ tmp.unlink(missing_ok=True)
87
+ raise DownloadError(f"HTTP {exc.code} for {url}") from exc
88
+ except urllib.error.URLError as exc:
89
+ tmp.unlink(missing_ok=True)
90
+ raise DownloadError(f"could not fetch {url}: {exc}") from exc
91
+ shutil.move(str(tmp), str(dest))
92
+ return dest
93
+
94
+
95
+ def _asset(name: str, filename: str, required: bool = True) -> Path | None:
96
+ dest = cache_dir() / name / filename
97
+ if dest.exists() and dest.stat().st_size:
98
+ return dest
99
+ url = f"{assets_url()}/{name}/{filename}"
100
+ try:
101
+ return download(url, dest)
102
+ except DownloadError:
103
+ if required:
104
+ raise
105
+ return None
106
+
107
+
108
+ def download_ir(name: str) -> Path:
109
+ """Fetch ``<name>.xml`` (+ ``.bin``, + ``labels.txt``); returns the .xml path."""
110
+ name = normalize_name(name)
111
+ if name not in MODEL_NAMES:
112
+ raise ModelNotFoundError(_unknown_name_message(name))
113
+ try:
114
+ xml = _asset(name, f"{name}.xml")
115
+ _asset(name, f"{name}.bin")
116
+ except DownloadError as exc:
117
+ raise ModelNotFoundError(
118
+ f"'{name}' is not on the mirror yet ({exc}). Pass a .pt/.xml path, "
119
+ f"point $EASYDETECT_ASSETS_URL at your own copy, or train it yourself: "
120
+ f"Detector('{name}').train(data='data.yaml')"
121
+ ) from exc
122
+ _asset(name, "labels.txt", required=False)
123
+ return xml
124
+
125
+
126
+ def download_checkpoint(name: str) -> Path:
127
+ """Fetch ``<name>.pt`` — the torch weights used to fine-tune or validate."""
128
+ name = normalize_name(name)
129
+ if name not in MODEL_NAMES:
130
+ raise ModelNotFoundError(_unknown_name_message(name))
131
+ try:
132
+ return _asset(name, f"{name}.pt")
133
+ except DownloadError as exc:
134
+ raise ModelNotFoundError(
135
+ f"no pretrained checkpoint for '{name}' on the mirror ({exc}). "
136
+ f"Training starts from an ImageNet backbone instead: "
137
+ f"Detector('{name}').train(data='data.yaml')"
138
+ ) from exc
139
+
140
+
141
+ def _unknown_name_message(name: str) -> str:
142
+ return (
143
+ f"unknown model '{name}'. Known names: {', '.join(MODEL_NAMES)}. "
144
+ f"You can also pass a path to your own .pt / .xml / .onnx."
145
+ )
easydetect/errors.py ADDED
@@ -0,0 +1,16 @@
1
+ # Apache-2.0
2
+ """Errors this package raises. Kept in one place so they are easy to catch."""
3
+
4
+ from __future__ import annotations
5
+
6
+
7
+ class EasyDetectError(Exception):
8
+ """Base class for every error raised by easydetect."""
9
+
10
+
11
+ class ModelNotFoundError(EasyDetectError, FileNotFoundError):
12
+ """A model name could not be resolved to weights (locally or on the mirror)."""
13
+
14
+
15
+ class DownloadError(EasyDetectError, OSError):
16
+ """A weight download failed."""
easydetect/exporter.py ADDED
@@ -0,0 +1,86 @@
1
+ # Apache-2.0
2
+ """Export: DFINENet -> ONNX (opset 17) -> OpenVINO IR, static shape, FP16 optional."""
3
+
4
+ from __future__ import annotations
5
+
6
+ import copy
7
+ import json
8
+ import threading
9
+ import warnings
10
+ from pathlib import Path
11
+
12
+ #: torch.onnx keeps global state while exporting ("in_onnx_export"), so two
13
+ #: exports at once fail an assertion deep inside it. Serialise them.
14
+ _EXPORT_LOCK = threading.Lock()
15
+
16
+
17
+ def _write_labels(out_dir: Path, fname: str, names) -> dict[int, str]:
18
+ """Write ``labels.txt`` (one class per line) + ``<fname>.names.json``.
19
+
20
+ labels.txt beside the IR is how downstream runtimes (ovkit among them)
21
+ discover class names, so a freshly trained model answers "my-class 0.91"
22
+ with no extra wiring.
23
+ """
24
+ table = {int(k): str(v) for k, v in (names or {}).items()}
25
+ (out_dir / f"{fname}.names.json").write_text(
26
+ json.dumps(table, ensure_ascii=False), encoding="utf-8"
27
+ )
28
+ if table:
29
+ lines = [table.get(i, f"class_{i}") for i in range(max(table) + 1)]
30
+ (out_dir / "labels.txt").write_text("\n".join(lines) + "\n", encoding="utf-8")
31
+ return table
32
+
33
+
34
+ def export_onnx(net, names, imgsz=640, out_dir=".", fname="easydetect", half=False, verbose=True):
35
+ """Write ``<out_dir>/<fname>.onnx`` (plus labels). Returns the .onnx path."""
36
+ import torch
37
+
38
+ from .nn import DeployWrapper
39
+
40
+ out_dir = Path(out_dir)
41
+ out_dir.mkdir(parents=True, exist_ok=True)
42
+ onnx_path = out_dir / f"{fname}.onnx"
43
+
44
+ # the deploy form fuses the re-parameterised blocks and drops the training
45
+ # heads; a copy, so the caller's network can keep training
46
+ deployed = copy.deepcopy(net).cpu()
47
+ deployed = deployed.deploy() if hasattr(deployed, "deploy") else deployed.eval()
48
+ wrapper = DeployWrapper(deployed).eval()
49
+ dummy = torch.zeros(1, 3, imgsz, imgsz)
50
+ with _EXPORT_LOCK, warnings.catch_warnings():
51
+ # The graph is exported at a fixed input size on purpose, so the tracer
52
+ # baking shape-dependent constants (the query count) in is what we want.
53
+ warnings.filterwarnings("ignore", category=torch.jit.TracerWarning)
54
+ warnings.filterwarnings("ignore", category=UserWarning) # shape-inference chatter
55
+ warnings.filterwarnings("ignore", category=DeprecationWarning)
56
+ torch.onnx.export(
57
+ wrapper,
58
+ dummy,
59
+ str(onnx_path),
60
+ input_names=["images"],
61
+ output_names=["boxes", "scores"],
62
+ opset_version=17,
63
+ dynamo=False,
64
+ )
65
+ _write_labels(out_dir, fname, names)
66
+ if verbose:
67
+ print(f"[easydetect] exported: {onnx_path}")
68
+ return onnx_path
69
+
70
+
71
+ def export_openvino(net, names, imgsz=640, out_dir=".", fname="easydetect", half=False,
72
+ verbose=True):
73
+ """Write ``<out_dir>/<fname>.xml`` (+ .bin, + labels). Returns the .xml path."""
74
+ import openvino as ov
75
+
76
+ out_dir = Path(out_dir)
77
+ onnx_path = export_onnx(
78
+ net, names, imgsz=imgsz, out_dir=out_dir, fname=fname, verbose=False
79
+ )
80
+ xml_path = out_dir / f"{fname}.xml"
81
+
82
+ model = ov.convert_model(str(onnx_path))
83
+ ov.save_model(model, str(xml_path), compress_to_fp16=half)
84
+ if verbose:
85
+ print(f"[easydetect] exported: {xml_path} ({'FP16' if half else 'FP32'})")
86
+ return xml_path