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 +44 -0
- easydetect/cli.py +149 -0
- easydetect/data/__init__.py +1 -0
- easydetect/data/dataset.py +161 -0
- easydetect/data/labels.py +48 -0
- easydetect/downloads.py +145 -0
- easydetect/errors.py +16 -0
- easydetect/exporter.py +86 -0
- easydetect/metrics.py +62 -0
- easydetect/model.py +606 -0
- easydetect/nn/__init__.py +4 -0
- easydetect/nn/decoder.py +962 -0
- easydetect/nn/denoising.py +125 -0
- easydetect/nn/dfine_net.py +106 -0
- easydetect/nn/encoder.py +483 -0
- easydetect/nn/hgnetv2.py +524 -0
- easydetect/nn/ops.py +407 -0
- easydetect/plotting.py +93 -0
- easydetect/predictor.py +170 -0
- easydetect/results.py +204 -0
- easydetect/sources.py +174 -0
- easydetect/tracker.py +87 -0
- easydetect/trainer.py +510 -0
- easydetect/utils/__init__.py +1 -0
- easydetect/utils/loss.py +686 -0
- easydetect/utils/ops.py +63 -0
- easydetect/validator.py +124 -0
- easydetect-0.1.0.dist-info/METADATA +250 -0
- easydetect-0.1.0.dist-info/RECORD +32 -0
- easydetect-0.1.0.dist-info/WHEEL +4 -0
- easydetect-0.1.0.dist-info/entry_points.txt +2 -0
- easydetect-0.1.0.dist-info/licenses/LICENSE +202 -0
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
|
easydetect/downloads.py
ADDED
|
@@ -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
|