kotoha 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.
- kotoha/__init__.py +163 -0
- kotoha/annotation.py +126 -0
- kotoha/common.py +421 -0
- kotoha/dataset_info.py +21 -0
- kotoha/eval/__init__.py +23 -0
- kotoha/eval/cls.py +465 -0
- kotoha/eval/det.py +486 -0
- kotoha/eval/pose.py +81 -0
- kotoha/eval/reid.py +216 -0
- kotoha/eval/text.py +357 -0
- kotoha/experiments/mcdm/__init__.py +17 -0
- kotoha/experiments/mcdm/_core.py +161 -0
- kotoha/experiments/mcdm/ranking.py +79 -0
- kotoha/experiments/mcdm/weights.py +113 -0
- kotoha/experiments/model/tracker/byte_track/__init__.py +3 -0
- kotoha/experiments/model/tracker/byte_track/byte_track.py +94 -0
- kotoha/experiments/model/tracker/byte_track/tracker/__init__.py +0 -0
- kotoha/experiments/model/tracker/byte_track/tracker/basetrack.py +53 -0
- kotoha/experiments/model/tracker/byte_track/tracker/byte_tracker.py +349 -0
- kotoha/experiments/model/tracker/byte_track/tracker/kalman_filter.py +259 -0
- kotoha/experiments/model/tracker/byte_track/tracker/matching.py +193 -0
- kotoha/experiments/model/tracker/oc_sort/__init__.py +3 -0
- kotoha/experiments/model/tracker/oc_sort/ocsort.py +27 -0
- kotoha/experiments/model/tracker/oc_sort/tracker/__init__.py +0 -0
- kotoha/experiments/model/tracker/oc_sort/tracker/association.py +381 -0
- kotoha/experiments/model/tracker/oc_sort/tracker/kalmanfilter.py +1533 -0
- kotoha/experiments/model/tracker/oc_sort/tracker/ocsort.py +343 -0
- kotoha/experiments/voc.py +409 -0
- kotoha/geo/__init__.py +17 -0
- kotoha/geo/_types.py +263 -0
- kotoha/geo/affine.py +51 -0
- kotoha/geo/circle.py +33 -0
- kotoha/geo/line.py +120 -0
- kotoha/geo/polygon.py +62 -0
- kotoha/geo/transform.py +29 -0
- kotoha/image_filter/__init__.py +3 -0
- kotoha/image_filter/cpp/__init__.py +3 -0
- kotoha/image_filter/hashing.py +165 -0
- kotoha/image_filter/image_filter.py +282 -0
- kotoha/io.py +449 -0
- kotoha/model/__init__.py +0 -0
- kotoha/model/mmlab/__init__.py +3 -0
- kotoha/model/mmlab/pose_heatmap.py +137 -0
- kotoha/model/mmlab/utils.py +54 -0
- kotoha/model/ultralytics/__init__.py +10 -0
- kotoha/model/ultralytics/base.py +142 -0
- kotoha/model/ultralytics/utils.py +314 -0
- kotoha/model/ultralytics/yolo26sem.py +114 -0
- kotoha/model/ultralytics/yolov5.py +110 -0
- kotoha/model/ultralytics/yolov8.py +92 -0
- kotoha/model/ultralytics/yolov8cls.py +138 -0
- kotoha/model/ultralytics/yolov8obb.py +126 -0
- kotoha/model/ultralytics/yolov8pose.py +134 -0
- kotoha/model/ultralytics/yolov8seg.py +106 -0
- kotoha/py.typed +1 -0
- kotoha/typing.py +43 -0
- kotoha/visual/__init__.py +20 -0
- kotoha/visual/color.py +74 -0
- kotoha/visual/visual.py +423 -0
- kotoha-0.1.0.dist-info/METADATA +206 -0
- kotoha-0.1.0.dist-info/RECORD +63 -0
- kotoha-0.1.0.dist-info/WHEEL +5 -0
- kotoha-0.1.0.dist-info/top_level.txt +1 -0
kotoha/__init__.py
ADDED
|
@@ -0,0 +1,163 @@
|
|
|
1
|
+
from .annotation import (
|
|
2
|
+
Classification,
|
|
3
|
+
Detection,
|
|
4
|
+
ObbDetection,
|
|
5
|
+
PoseDetection,
|
|
6
|
+
SegmentationDetection,
|
|
7
|
+
SemanticSegmentation,
|
|
8
|
+
)
|
|
9
|
+
from .common import (
|
|
10
|
+
FuzzyMatchingSet,
|
|
11
|
+
is_url,
|
|
12
|
+
is_valid_image,
|
|
13
|
+
parallel_process,
|
|
14
|
+
scale_bbox_xyxy,
|
|
15
|
+
url_to_image,
|
|
16
|
+
)
|
|
17
|
+
from .dataset_info import CocoConfig
|
|
18
|
+
from .eval import (
|
|
19
|
+
average_precision,
|
|
20
|
+
bleu,
|
|
21
|
+
confusion,
|
|
22
|
+
edit_distance,
|
|
23
|
+
evaluate_detection,
|
|
24
|
+
evaluate_reid_market1501,
|
|
25
|
+
evaluate_reid_roc,
|
|
26
|
+
f1_score,
|
|
27
|
+
match_detections,
|
|
28
|
+
mean_average_precision,
|
|
29
|
+
pairwise_box_iou,
|
|
30
|
+
pck_accuracy,
|
|
31
|
+
rouge,
|
|
32
|
+
safe_divide,
|
|
33
|
+
topk_accuracy,
|
|
34
|
+
)
|
|
35
|
+
from .geo import (
|
|
36
|
+
circle_from_three_points,
|
|
37
|
+
fit_line,
|
|
38
|
+
intersect_lines,
|
|
39
|
+
is_point_in_polygon,
|
|
40
|
+
is_point_on_segment,
|
|
41
|
+
polygon_area,
|
|
42
|
+
rotate_points,
|
|
43
|
+
segment_intersection,
|
|
44
|
+
)
|
|
45
|
+
from .image_filter import ImageFilter
|
|
46
|
+
from .io import (
|
|
47
|
+
list_files,
|
|
48
|
+
read_csv,
|
|
49
|
+
read_json,
|
|
50
|
+
read_pkl,
|
|
51
|
+
read_txt,
|
|
52
|
+
read_xml,
|
|
53
|
+
read_yaml,
|
|
54
|
+
read_yolo_txt,
|
|
55
|
+
save_csv,
|
|
56
|
+
save_json,
|
|
57
|
+
save_pkl,
|
|
58
|
+
save_txt,
|
|
59
|
+
save_xml,
|
|
60
|
+
save_yaml,
|
|
61
|
+
save_yolo_txt,
|
|
62
|
+
)
|
|
63
|
+
from .typing import (
|
|
64
|
+
Affine2D,
|
|
65
|
+
Circle,
|
|
66
|
+
KeyPoint,
|
|
67
|
+
Line,
|
|
68
|
+
Point2f,
|
|
69
|
+
PointGeometryMixin,
|
|
70
|
+
PointLike,
|
|
71
|
+
PointSet,
|
|
72
|
+
Polygon,
|
|
73
|
+
PolygonLike,
|
|
74
|
+
Rect,
|
|
75
|
+
Vector2,
|
|
76
|
+
)
|
|
77
|
+
from .visual import (
|
|
78
|
+
VideoReader,
|
|
79
|
+
draw_bbox,
|
|
80
|
+
draw_keypoints,
|
|
81
|
+
draw_masks,
|
|
82
|
+
generate_distinct_colors,
|
|
83
|
+
get_color,
|
|
84
|
+
imread,
|
|
85
|
+
imwrite,
|
|
86
|
+
)
|
|
87
|
+
|
|
88
|
+
__all__ = [
|
|
89
|
+
"Affine2D",
|
|
90
|
+
"Circle",
|
|
91
|
+
"Classification",
|
|
92
|
+
"CocoConfig",
|
|
93
|
+
"Detection",
|
|
94
|
+
"FuzzyMatchingSet",
|
|
95
|
+
"ImageFilter",
|
|
96
|
+
"KeyPoint",
|
|
97
|
+
"Line",
|
|
98
|
+
"ObbDetection",
|
|
99
|
+
"Point2f",
|
|
100
|
+
"PointGeometryMixin",
|
|
101
|
+
"PointLike",
|
|
102
|
+
"PointSet",
|
|
103
|
+
"Polygon",
|
|
104
|
+
"PolygonLike",
|
|
105
|
+
"PoseDetection",
|
|
106
|
+
"Rect",
|
|
107
|
+
"SegmentationDetection",
|
|
108
|
+
"SemanticSegmentation",
|
|
109
|
+
"Vector2",
|
|
110
|
+
"VideoReader",
|
|
111
|
+
"average_precision",
|
|
112
|
+
"bleu",
|
|
113
|
+
"circle_from_three_points",
|
|
114
|
+
"confusion",
|
|
115
|
+
"draw_bbox",
|
|
116
|
+
"draw_keypoints",
|
|
117
|
+
"draw_masks",
|
|
118
|
+
"edit_distance",
|
|
119
|
+
"evaluate_detection",
|
|
120
|
+
"evaluate_reid_market1501",
|
|
121
|
+
"evaluate_reid_roc",
|
|
122
|
+
"f1_score",
|
|
123
|
+
"fit_line",
|
|
124
|
+
"generate_distinct_colors",
|
|
125
|
+
"get_color",
|
|
126
|
+
"imread",
|
|
127
|
+
"imwrite",
|
|
128
|
+
"intersect_lines",
|
|
129
|
+
"is_point_in_polygon",
|
|
130
|
+
"is_point_on_segment",
|
|
131
|
+
"is_url",
|
|
132
|
+
"is_valid_image",
|
|
133
|
+
"list_files",
|
|
134
|
+
"match_detections",
|
|
135
|
+
"mean_average_precision",
|
|
136
|
+
"pairwise_box_iou",
|
|
137
|
+
"parallel_process",
|
|
138
|
+
"pck_accuracy",
|
|
139
|
+
"polygon_area",
|
|
140
|
+
"read_csv",
|
|
141
|
+
"read_json",
|
|
142
|
+
"read_pkl",
|
|
143
|
+
"read_txt",
|
|
144
|
+
"read_xml",
|
|
145
|
+
"read_yaml",
|
|
146
|
+
"read_yolo_txt",
|
|
147
|
+
"rotate_points",
|
|
148
|
+
"rouge",
|
|
149
|
+
"safe_divide",
|
|
150
|
+
"save_csv",
|
|
151
|
+
"save_json",
|
|
152
|
+
"save_pkl",
|
|
153
|
+
"save_txt",
|
|
154
|
+
"save_xml",
|
|
155
|
+
"save_yaml",
|
|
156
|
+
"save_yolo_txt",
|
|
157
|
+
"scale_bbox_xyxy",
|
|
158
|
+
"segment_intersection",
|
|
159
|
+
"topk_accuracy",
|
|
160
|
+
"url_to_image",
|
|
161
|
+
]
|
|
162
|
+
|
|
163
|
+
__version__ = "0.1.0"
|
kotoha/annotation.py
ADDED
|
@@ -0,0 +1,126 @@
|
|
|
1
|
+
from __future__ import annotations
|
|
2
|
+
|
|
3
|
+
from dataclasses import asdict, dataclass
|
|
4
|
+
|
|
5
|
+
import numpy as np
|
|
6
|
+
|
|
7
|
+
from kotoha.geo._types import KeyPoint, Polygon, Rect
|
|
8
|
+
|
|
9
|
+
__all__ = [
|
|
10
|
+
"Classification",
|
|
11
|
+
"Detection",
|
|
12
|
+
"ObbDetection",
|
|
13
|
+
"PoseDetection",
|
|
14
|
+
"SegmentationDetection",
|
|
15
|
+
"SemanticSegmentation",
|
|
16
|
+
]
|
|
17
|
+
|
|
18
|
+
|
|
19
|
+
@dataclass(slots=True)
|
|
20
|
+
class Classification:
|
|
21
|
+
"""One ranked category prediction for an image."""
|
|
22
|
+
|
|
23
|
+
score: float
|
|
24
|
+
class_id: int
|
|
25
|
+
label: str | None = None
|
|
26
|
+
|
|
27
|
+
def to_numpy(self) -> np.ndarray:
|
|
28
|
+
return np.array([self.score, self.class_id], dtype=np.float32)
|
|
29
|
+
|
|
30
|
+
def to_dict(self) -> dict[str, object]:
|
|
31
|
+
return asdict(self)
|
|
32
|
+
|
|
33
|
+
|
|
34
|
+
@dataclass(slots=True)
|
|
35
|
+
class SemanticSegmentation:
|
|
36
|
+
"""Dense semantic class map for one image."""
|
|
37
|
+
|
|
38
|
+
mask: np.ndarray
|
|
39
|
+
|
|
40
|
+
def to_numpy(self) -> np.ndarray:
|
|
41
|
+
return self.mask
|
|
42
|
+
|
|
43
|
+
def to_dict(self) -> dict[str, object]:
|
|
44
|
+
return {"mask": self.mask.tolist()}
|
|
45
|
+
|
|
46
|
+
|
|
47
|
+
@dataclass(slots=True)
|
|
48
|
+
class Detection:
|
|
49
|
+
box: Rect
|
|
50
|
+
score: float
|
|
51
|
+
class_id: int
|
|
52
|
+
label: str | None = None
|
|
53
|
+
|
|
54
|
+
def to_numpy(self) -> np.ndarray:
|
|
55
|
+
return np.array(
|
|
56
|
+
[*self.box.to_list(), self.score, self.class_id],
|
|
57
|
+
dtype=np.float32,
|
|
58
|
+
)
|
|
59
|
+
|
|
60
|
+
def to_dict(self) -> dict[str, object]:
|
|
61
|
+
data = asdict(self)
|
|
62
|
+
data["box"] = self.box.to_dict()
|
|
63
|
+
return data
|
|
64
|
+
|
|
65
|
+
|
|
66
|
+
@dataclass(slots=True)
|
|
67
|
+
class SegmentationDetection:
|
|
68
|
+
box: Rect
|
|
69
|
+
score: float
|
|
70
|
+
class_id: int
|
|
71
|
+
mask: np.ndarray
|
|
72
|
+
polygon: Polygon | None = None
|
|
73
|
+
label: str | None = None
|
|
74
|
+
|
|
75
|
+
def to_numpy(self) -> np.ndarray:
|
|
76
|
+
return np.array(
|
|
77
|
+
[*self.box.to_list(), self.score, self.class_id],
|
|
78
|
+
dtype=np.float32,
|
|
79
|
+
)
|
|
80
|
+
|
|
81
|
+
def to_dict(self) -> dict[str, object]:
|
|
82
|
+
data = asdict(self)
|
|
83
|
+
data["box"] = self.box.to_dict()
|
|
84
|
+
data["mask"] = self.mask.tolist()
|
|
85
|
+
if self.polygon is not None:
|
|
86
|
+
data["polygon"] = self.polygon.to_dict()
|
|
87
|
+
return data
|
|
88
|
+
|
|
89
|
+
|
|
90
|
+
@dataclass(slots=True)
|
|
91
|
+
class PoseDetection:
|
|
92
|
+
box: Rect
|
|
93
|
+
score: float
|
|
94
|
+
class_id: int
|
|
95
|
+
keypoints: list[KeyPoint]
|
|
96
|
+
label: str | None = None
|
|
97
|
+
|
|
98
|
+
def to_numpy(self) -> np.ndarray:
|
|
99
|
+
values: list[float] = [*self.box.to_list(), self.score, float(self.class_id)]
|
|
100
|
+
for keypoint in self.keypoints:
|
|
101
|
+
values.extend([keypoint.p.x, keypoint.p.y, keypoint.score or 0.0])
|
|
102
|
+
return np.array(values, dtype=np.float32)
|
|
103
|
+
|
|
104
|
+
def to_dict(self) -> dict[str, object]:
|
|
105
|
+
data = asdict(self)
|
|
106
|
+
data["keypoints"] = [keypoint.to_dict() for keypoint in self.keypoints]
|
|
107
|
+
data["box"] = self.box.to_dict()
|
|
108
|
+
return data
|
|
109
|
+
|
|
110
|
+
|
|
111
|
+
@dataclass(slots=True)
|
|
112
|
+
class ObbDetection:
|
|
113
|
+
quad: Polygon
|
|
114
|
+
score: float
|
|
115
|
+
class_id: int
|
|
116
|
+
label: str | None = None
|
|
117
|
+
|
|
118
|
+
def to_numpy(self) -> np.ndarray:
|
|
119
|
+
values = [coord for point in self.quad.vertices for coord in (point.x, point.y)]
|
|
120
|
+
values.extend([self.score, float(self.class_id)])
|
|
121
|
+
return np.array(values, dtype=np.float32)
|
|
122
|
+
|
|
123
|
+
def to_dict(self) -> dict[str, object]:
|
|
124
|
+
data = asdict(self)
|
|
125
|
+
data["quad"] = self.quad.to_dict()
|
|
126
|
+
return data
|
kotoha/common.py
ADDED
|
@@ -0,0 +1,421 @@
|
|
|
1
|
+
import logging
|
|
2
|
+
import time
|
|
3
|
+
from collections import Counter, defaultdict
|
|
4
|
+
from collections.abc import Callable, Iterable, Sequence, Sized
|
|
5
|
+
from multiprocessing import Pool as ProcessPool
|
|
6
|
+
from multiprocessing.pool import ThreadPool
|
|
7
|
+
from pathlib import Path
|
|
8
|
+
from typing import Any, Literal, overload
|
|
9
|
+
from urllib import parse, request
|
|
10
|
+
|
|
11
|
+
import cv2
|
|
12
|
+
import numpy as np
|
|
13
|
+
import requests
|
|
14
|
+
from PIL import Image
|
|
15
|
+
from tqdm import tqdm
|
|
16
|
+
|
|
17
|
+
from kotoha.eval import edit_distance
|
|
18
|
+
|
|
19
|
+
logger = logging.getLogger(__name__)
|
|
20
|
+
|
|
21
|
+
__all__ = [
|
|
22
|
+
"FuzzyMatchingSet",
|
|
23
|
+
"is_url",
|
|
24
|
+
"is_valid_image",
|
|
25
|
+
"parallel_process",
|
|
26
|
+
"scale_bbox_xyxy",
|
|
27
|
+
"url_to_image",
|
|
28
|
+
]
|
|
29
|
+
|
|
30
|
+
|
|
31
|
+
def is_url(url: str, check: bool = False) -> bool:
|
|
32
|
+
"""
|
|
33
|
+
Validate if the given string is a URL and optionally check if the URL exists online.
|
|
34
|
+
|
|
35
|
+
Args:
|
|
36
|
+
url (str): The string to be validated as a URL.
|
|
37
|
+
check (bool, optional): If True, performs an additional check to see if the URL exists online.
|
|
38
|
+
|
|
39
|
+
Returns:
|
|
40
|
+
(bool): True for a valid URL. If 'check' is True, also returns True if the URL exists online.
|
|
41
|
+
|
|
42
|
+
Examples:
|
|
43
|
+
>>> valid = is_url("https://www.example.com")
|
|
44
|
+
>>> valid_and_exists = is_url("https://www.example.com", check=True)
|
|
45
|
+
"""
|
|
46
|
+
try:
|
|
47
|
+
url = str(url)
|
|
48
|
+
result = parse.urlparse(url)
|
|
49
|
+
assert all([result.scheme, result.netloc]) # check if is url
|
|
50
|
+
if check:
|
|
51
|
+
with request.urlopen(url) as response:
|
|
52
|
+
return response.getcode() == 200 # check if exists online
|
|
53
|
+
return True
|
|
54
|
+
except Exception:
|
|
55
|
+
return False
|
|
56
|
+
|
|
57
|
+
|
|
58
|
+
def url_to_image(url: str, readFlag: int = cv2.IMREAD_COLOR, headers=None) -> np.ndarray | None:
|
|
59
|
+
"""
|
|
60
|
+
Download an image from a URL and decode it into an OpenCV image.
|
|
61
|
+
|
|
62
|
+
Args:
|
|
63
|
+
url (str): URL of the image to download.
|
|
64
|
+
readFlag (int, optional): Flag specifying the color type of a loaded image.
|
|
65
|
+
Defaults to cv2.IMREAD_COLOR.
|
|
66
|
+
|
|
67
|
+
Returns:
|
|
68
|
+
Optional[np.ndarray]: Decoded image as a numpy array if successful, else None.
|
|
69
|
+
"""
|
|
70
|
+
if headers is None:
|
|
71
|
+
headers = {
|
|
72
|
+
"User-Agent": "Mozilla/5.0 (Windows NT 10.0; Win64; x64) AppleWebKit/537.36 (KHTML, like Gecko) Chrome/91.0.4472.124 Safari/537.36"
|
|
73
|
+
}
|
|
74
|
+
try:
|
|
75
|
+
response = requests.get(url, headers=headers, timeout=10)
|
|
76
|
+
response.raise_for_status()
|
|
77
|
+
image_array = np.frombuffer(response.content, dtype=np.uint8)
|
|
78
|
+
image = cv2.imdecode(image_array, readFlag)
|
|
79
|
+
return image
|
|
80
|
+
except Exception as e:
|
|
81
|
+
logger.warning("Failed to download or decode image from %s: %s", url, e)
|
|
82
|
+
return None
|
|
83
|
+
|
|
84
|
+
|
|
85
|
+
def is_valid_image(path: str | Path) -> bool:
|
|
86
|
+
"""
|
|
87
|
+
Checks whether the given file is a valid image by attempting to open and verify it.
|
|
88
|
+
|
|
89
|
+
Args:
|
|
90
|
+
path (str | Path): Path to the image file.
|
|
91
|
+
|
|
92
|
+
Returns:
|
|
93
|
+
bool: True if the image is valid, False otherwise.
|
|
94
|
+
|
|
95
|
+
Raises:
|
|
96
|
+
None: All exceptions are caught internally and False is returned.
|
|
97
|
+
"""
|
|
98
|
+
try:
|
|
99
|
+
with Image.open(path) as img:
|
|
100
|
+
img.verify() # Verify that it is, in fact, an image
|
|
101
|
+
return True
|
|
102
|
+
except Exception:
|
|
103
|
+
return False
|
|
104
|
+
|
|
105
|
+
|
|
106
|
+
@overload
|
|
107
|
+
def parallel_process[T](
|
|
108
|
+
func: Callable[[Any], T],
|
|
109
|
+
data: Iterable,
|
|
110
|
+
*,
|
|
111
|
+
use_threads: bool = False,
|
|
112
|
+
num_workers: int = 4,
|
|
113
|
+
store_results: Literal[True] = True,
|
|
114
|
+
show_progress: bool = True,
|
|
115
|
+
prog_desc: str | None = None,
|
|
116
|
+
prog_leave: bool = True,
|
|
117
|
+
) -> tuple[list[T], float]: ...
|
|
118
|
+
|
|
119
|
+
|
|
120
|
+
@overload
|
|
121
|
+
def parallel_process[T](
|
|
122
|
+
func: Callable[[Any], T],
|
|
123
|
+
data: Iterable,
|
|
124
|
+
*,
|
|
125
|
+
use_threads: bool = False,
|
|
126
|
+
num_workers: int = 4,
|
|
127
|
+
store_results: Literal[False],
|
|
128
|
+
show_progress: bool = True,
|
|
129
|
+
prog_desc: str | None = None,
|
|
130
|
+
prog_leave: bool = True,
|
|
131
|
+
) -> tuple[None, float]: ...
|
|
132
|
+
|
|
133
|
+
|
|
134
|
+
def parallel_process[T](
|
|
135
|
+
func: Callable[[Any], T],
|
|
136
|
+
data: Iterable,
|
|
137
|
+
*,
|
|
138
|
+
use_threads: bool = False,
|
|
139
|
+
num_workers: int = 4,
|
|
140
|
+
store_results: bool = True,
|
|
141
|
+
show_progress: bool = True,
|
|
142
|
+
prog_desc: str | None = None,
|
|
143
|
+
prog_leave: bool = True,
|
|
144
|
+
):
|
|
145
|
+
"""
|
|
146
|
+
General-purpose function for parallel processing using multiple processes.
|
|
147
|
+
|
|
148
|
+
Args:
|
|
149
|
+
func (Callable):
|
|
150
|
+
The function to execute in parallel. It should accept a single input item
|
|
151
|
+
(each element of `data`) as its argument.
|
|
152
|
+
|
|
153
|
+
data (Iterable):
|
|
154
|
+
An iterable of input items, each passed as the argument to `func`.
|
|
155
|
+
Can be any iterable — e.g. list, generator, or range.
|
|
156
|
+
|
|
157
|
+
use_threads (bool, optional):
|
|
158
|
+
If True, use ThreadPool. Defaults to False (ProcessPool).
|
|
159
|
+
|
|
160
|
+
num_workers (int, optional):
|
|
161
|
+
Number of worker processes to use. Defaults to 4.
|
|
162
|
+
|
|
163
|
+
store_results (bool, optional):
|
|
164
|
+
Whether to collect and return the function outputs.
|
|
165
|
+
If False, results are discarded after execution (useful for side-effect-only tasks).
|
|
166
|
+
Defaults to True.
|
|
167
|
+
|
|
168
|
+
show_progress (bool): Whether to display a progress bar during processing.
|
|
169
|
+
|
|
170
|
+
prog_desc (str, optional):
|
|
171
|
+
Custom description text for the progress bar (from tqdm). Defaults to None.
|
|
172
|
+
|
|
173
|
+
prog_leave (bool, optional):
|
|
174
|
+
Whether to leave the progress bar on screen after completion. Defaults to True.
|
|
175
|
+
|
|
176
|
+
Returns:
|
|
177
|
+
Tuple[List[Any] | None, float]:
|
|
178
|
+
- results: List of all outputs from `func`, if `store_results=True`. Otherwise, None.
|
|
179
|
+
- duration: Total wall-clock execution time in seconds.
|
|
180
|
+
|
|
181
|
+
Raises:
|
|
182
|
+
TypeError:
|
|
183
|
+
If `func` is not callable.
|
|
184
|
+
|
|
185
|
+
Note:
|
|
186
|
+
Exceptions raised inside child processes are caught and logged internally.
|
|
187
|
+
The main process will not crash due to worker errors.
|
|
188
|
+
"""
|
|
189
|
+
|
|
190
|
+
if not callable(func):
|
|
191
|
+
raise TypeError("func must be a callable function.")
|
|
192
|
+
|
|
193
|
+
start_time = time.time()
|
|
194
|
+
|
|
195
|
+
results = [] if store_results else None
|
|
196
|
+
|
|
197
|
+
PoolClass = ThreadPool if use_threads else ProcessPool
|
|
198
|
+
|
|
199
|
+
total_items = len(data) if isinstance(data, Sized) else None
|
|
200
|
+
|
|
201
|
+
with PoolClass(processes=num_workers) as pool:
|
|
202
|
+
iterator = pool.imap_unordered(func, data)
|
|
203
|
+
if show_progress:
|
|
204
|
+
iterator = tqdm(iterator, desc=prog_desc, leave=prog_leave, total=total_items)
|
|
205
|
+
for res in iterator:
|
|
206
|
+
if results is not None and res is not None:
|
|
207
|
+
results.append(res)
|
|
208
|
+
|
|
209
|
+
duration = time.time() - start_time
|
|
210
|
+
|
|
211
|
+
return (results, duration)
|
|
212
|
+
|
|
213
|
+
|
|
214
|
+
class FuzzyMatchingSet:
|
|
215
|
+
"""
|
|
216
|
+
Fuzzy string matcher based on n-gram inverted index and edit distance.
|
|
217
|
+
|
|
218
|
+
The matcher first uses n-gram features to collect a small set of candidate
|
|
219
|
+
words, then ranks those candidates by normalized edit distance.
|
|
220
|
+
|
|
221
|
+
Args:
|
|
222
|
+
words: Candidate words.
|
|
223
|
+
ngram_size: Feature width used for indexing. ``2`` is usually a good
|
|
224
|
+
default for short labels.
|
|
225
|
+
max_extra_length: Reject query words that are much longer than all
|
|
226
|
+
indexed words.
|
|
227
|
+
max_candidates: Maximum number of candidates to verify with edit distance.
|
|
228
|
+
case_sensitive: Whether matching should be case-sensitive.
|
|
229
|
+
"""
|
|
230
|
+
|
|
231
|
+
def __init__(
|
|
232
|
+
self,
|
|
233
|
+
words: Iterable[str],
|
|
234
|
+
*,
|
|
235
|
+
ngram_size: int = 2,
|
|
236
|
+
max_extra_length: int = 2,
|
|
237
|
+
max_candidates: int = 50,
|
|
238
|
+
case_sensitive: bool = False,
|
|
239
|
+
):
|
|
240
|
+
if ngram_size <= 0:
|
|
241
|
+
raise ValueError("ngram_size must be a positive integer")
|
|
242
|
+
if max_extra_length < 0:
|
|
243
|
+
raise ValueError("max_extra_length must be non-negative")
|
|
244
|
+
if max_candidates <= 0:
|
|
245
|
+
raise ValueError("max_candidates must be a positive integer")
|
|
246
|
+
|
|
247
|
+
self.ngram_size = ngram_size
|
|
248
|
+
self.max_extra_length = max_extra_length
|
|
249
|
+
self.max_candidates = max_candidates
|
|
250
|
+
self.case_sensitive = case_sensitive
|
|
251
|
+
|
|
252
|
+
self.words = list(dict.fromkeys(words))
|
|
253
|
+
self.normalized_words = [self._normalize(word) for word in self.words]
|
|
254
|
+
|
|
255
|
+
self.max_word_len = max((len(word) for word in self.normalized_words), default=0)
|
|
256
|
+
|
|
257
|
+
self.direct_indexes: dict[str, int] = {}
|
|
258
|
+
self.feature_indexes: dict[str, set[int]] = defaultdict(set)
|
|
259
|
+
|
|
260
|
+
for index, word in enumerate(self.normalized_words):
|
|
261
|
+
self.direct_indexes[word] = index
|
|
262
|
+
|
|
263
|
+
for feature in self._get_features(word):
|
|
264
|
+
self.feature_indexes[feature].add(index)
|
|
265
|
+
|
|
266
|
+
def match(self, word: str, threshold: float = 0.3) -> str | None:
|
|
267
|
+
"""
|
|
268
|
+
Find the closest word in the set.
|
|
269
|
+
|
|
270
|
+
Args:
|
|
271
|
+
word: Query word.
|
|
272
|
+
threshold: Maximum normalized edit distance allowed. Smaller values
|
|
273
|
+
are stricter. For example, ``0.3`` means the edit distance must
|
|
274
|
+
be no more than 30% of the longer word length.
|
|
275
|
+
|
|
276
|
+
Returns:
|
|
277
|
+
The best matched original word, or ``None`` if no candidate is good enough.
|
|
278
|
+
"""
|
|
279
|
+
result = self.match_with_score(word, threshold=threshold)
|
|
280
|
+
if result is None:
|
|
281
|
+
return None
|
|
282
|
+
|
|
283
|
+
matched_word, _ = result
|
|
284
|
+
return matched_word
|
|
285
|
+
|
|
286
|
+
def match_with_score(
|
|
287
|
+
self,
|
|
288
|
+
word: str,
|
|
289
|
+
threshold: float = 0.3,
|
|
290
|
+
) -> tuple[str, float] | None:
|
|
291
|
+
"""
|
|
292
|
+
Find the closest word and return its normalized distance.
|
|
293
|
+
|
|
294
|
+
Args:
|
|
295
|
+
word: Query word.
|
|
296
|
+
threshold: Maximum normalized edit distance allowed.
|
|
297
|
+
|
|
298
|
+
Returns:
|
|
299
|
+
``(matched_word, score)`` if matched, otherwise ``None``.
|
|
300
|
+
The score is normalized edit distance, so smaller is better.
|
|
301
|
+
"""
|
|
302
|
+
if threshold < 0:
|
|
303
|
+
raise ValueError("threshold must be non-negative")
|
|
304
|
+
|
|
305
|
+
if not self.words:
|
|
306
|
+
return None
|
|
307
|
+
|
|
308
|
+
query = self._normalize(word)
|
|
309
|
+
|
|
310
|
+
if len(query) > self.max_word_len + self.max_extra_length:
|
|
311
|
+
return None
|
|
312
|
+
|
|
313
|
+
direct_index = self.direct_indexes.get(query)
|
|
314
|
+
if direct_index is not None:
|
|
315
|
+
return self.words[direct_index], 0.0
|
|
316
|
+
|
|
317
|
+
candidates = self._get_candidates(query)
|
|
318
|
+
if not candidates:
|
|
319
|
+
return None
|
|
320
|
+
|
|
321
|
+
best_index: int | None = None
|
|
322
|
+
best_score: float | None = None
|
|
323
|
+
|
|
324
|
+
for index in candidates:
|
|
325
|
+
candidate = self.normalized_words[index]
|
|
326
|
+
score = edit_distance(query, candidate) / max(len(query), len(candidate), 1)
|
|
327
|
+
|
|
328
|
+
if best_score is None or score < best_score:
|
|
329
|
+
best_index = index
|
|
330
|
+
best_score = score
|
|
331
|
+
|
|
332
|
+
if best_index is None or best_score is None:
|
|
333
|
+
return None
|
|
334
|
+
|
|
335
|
+
if best_score <= threshold:
|
|
336
|
+
return self.words[best_index], best_score
|
|
337
|
+
|
|
338
|
+
return None
|
|
339
|
+
|
|
340
|
+
def _get_candidates(self, word: str) -> list[int]:
|
|
341
|
+
features = self._get_features(word)
|
|
342
|
+
counter: Counter[int] = Counter()
|
|
343
|
+
|
|
344
|
+
for feature in features:
|
|
345
|
+
for index in self.feature_indexes.get(feature, ()):
|
|
346
|
+
counter[index] += 1
|
|
347
|
+
|
|
348
|
+
return [index for index, _ in counter.most_common(self.max_candidates)]
|
|
349
|
+
|
|
350
|
+
def _normalize(self, word: str) -> str:
|
|
351
|
+
if self.case_sensitive:
|
|
352
|
+
return word
|
|
353
|
+
return word.lower()
|
|
354
|
+
|
|
355
|
+
def _get_features(self, word: str) -> tuple[str, ...]:
|
|
356
|
+
if not word:
|
|
357
|
+
return ("",)
|
|
358
|
+
|
|
359
|
+
if len(word) <= self.ngram_size:
|
|
360
|
+
return (word,)
|
|
361
|
+
|
|
362
|
+
return tuple(word[i : i + self.ngram_size] for i in range(len(word) - self.ngram_size + 1))
|
|
363
|
+
|
|
364
|
+
|
|
365
|
+
def scale_bbox_xyxy(
|
|
366
|
+
xyxy: Sequence[float],
|
|
367
|
+
scale: float | tuple[float, float] = 1.0,
|
|
368
|
+
image_w: int | None = None,
|
|
369
|
+
image_h: int | None = None,
|
|
370
|
+
min_size: float = 1.0,
|
|
371
|
+
) -> np.ndarray:
|
|
372
|
+
"""
|
|
373
|
+
Scale (expand/shrink) a bbox centered at its midpoint and optionally clip to image bounds.
|
|
374
|
+
|
|
375
|
+
Args:
|
|
376
|
+
xyxy: [x1, y1, x2, y2]
|
|
377
|
+
scale: float or (sx, sy)
|
|
378
|
+
- float: uniform scaling
|
|
379
|
+
- tuple: separate x/y scaling
|
|
380
|
+
image_w: optional image width for clipping
|
|
381
|
+
image_h: optional image height for clipping
|
|
382
|
+
min_size: minimum width/height to avoid degenerate boxes
|
|
383
|
+
|
|
384
|
+
Returns:
|
|
385
|
+
np.ndarray: [x1, y1, x2, y2] float32
|
|
386
|
+
"""
|
|
387
|
+
x1, y1, x2, y2 = map(float, xyxy)
|
|
388
|
+
|
|
389
|
+
# fix invalid box
|
|
390
|
+
if x2 < x1:
|
|
391
|
+
x1, x2 = x2, x1
|
|
392
|
+
if y2 < y1:
|
|
393
|
+
y1, y2 = y2, y1
|
|
394
|
+
|
|
395
|
+
sx, sy = (scale, scale) if isinstance(scale, (float, int)) else scale
|
|
396
|
+
|
|
397
|
+
cx = (x1 + x2) * 0.5
|
|
398
|
+
cy = (y1 + y2) * 0.5
|
|
399
|
+
|
|
400
|
+
bw = max((x2 - x1) * sx, min_size)
|
|
401
|
+
bh = max((y2 - y1) * sy, min_size)
|
|
402
|
+
|
|
403
|
+
nx1 = cx - bw * 0.5
|
|
404
|
+
ny1 = cy - bh * 0.5
|
|
405
|
+
nx2 = cx + bw * 0.5
|
|
406
|
+
ny2 = cy + bh * 0.5
|
|
407
|
+
|
|
408
|
+
# clip (careful: image boundary is usually [0, w-1])
|
|
409
|
+
if image_w is not None:
|
|
410
|
+
nx1 = np.clip(nx1, 0, image_w - 1)
|
|
411
|
+
nx2 = np.clip(nx2, 0, image_w - 1)
|
|
412
|
+
|
|
413
|
+
if image_h is not None:
|
|
414
|
+
ny1 = np.clip(ny1, 0, image_h - 1)
|
|
415
|
+
ny2 = np.clip(ny2, 0, image_h - 1)
|
|
416
|
+
|
|
417
|
+
# re-ensure ordering after clipping
|
|
418
|
+
x1, x2 = sorted((nx1, nx2))
|
|
419
|
+
y1, y2 = sorted((ny1, ny2))
|
|
420
|
+
|
|
421
|
+
return np.array([x1, y1, x2, y2], dtype=np.float32)
|