cvat-cli 2.68.0__tar.gz → 2.69.0__tar.gz
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.
- {cvat_cli-2.68.0 → cvat_cli-2.69.0}/PKG-INFO +2 -2
- {cvat_cli-2.68.0 → cvat_cli-2.69.0}/pyproject.toml +1 -1
- {cvat_cli-2.68.0 → cvat_cli-2.69.0}/src/cvat_cli/_internal/agent.py +55 -405
- cvat_cli-2.69.0/src/cvat_cli/_internal/agent_driver.py +95 -0
- cvat_cli-2.69.0/src/cvat_cli/_internal/agent_driver_detection.py +167 -0
- cvat_cli-2.69.0/src/cvat_cli/_internal/agent_driver_tracking.py +214 -0
- {cvat_cli-2.68.0 → cvat_cli-2.69.0}/src/cvat_cli/_internal/commands_functions.py +5 -8
- cvat_cli-2.69.0/src/cvat_cli/version.py +1 -0
- {cvat_cli-2.68.0 → cvat_cli-2.69.0}/src/cvat_cli.egg-info/PKG-INFO +2 -2
- {cvat_cli-2.68.0 → cvat_cli-2.69.0}/src/cvat_cli.egg-info/SOURCES.txt +3 -0
- {cvat_cli-2.68.0 → cvat_cli-2.69.0}/src/cvat_cli.egg-info/requires.txt +1 -1
- cvat_cli-2.68.0/src/cvat_cli/version.py +0 -1
- {cvat_cli-2.68.0 → cvat_cli-2.69.0}/README.md +0 -0
- {cvat_cli-2.68.0 → cvat_cli-2.69.0}/setup.cfg +0 -0
- {cvat_cli-2.68.0 → cvat_cli-2.69.0}/src/cvat_cli/__init__.py +0 -0
- {cvat_cli-2.68.0 → cvat_cli-2.69.0}/src/cvat_cli/__main__.py +0 -0
- {cvat_cli-2.68.0 → cvat_cli-2.69.0}/src/cvat_cli/_internal/__init__.py +0 -0
- {cvat_cli-2.68.0 → cvat_cli-2.69.0}/src/cvat_cli/_internal/command_base.py +0 -0
- {cvat_cli-2.68.0 → cvat_cli-2.69.0}/src/cvat_cli/_internal/commands_all.py +0 -0
- {cvat_cli-2.68.0 → cvat_cli-2.69.0}/src/cvat_cli/_internal/commands_projects.py +0 -0
- {cvat_cli-2.68.0 → cvat_cli-2.69.0}/src/cvat_cli/_internal/commands_tasks.py +0 -0
- {cvat_cli-2.68.0 → cvat_cli-2.69.0}/src/cvat_cli/_internal/common.py +0 -0
- {cvat_cli-2.68.0 → cvat_cli-2.69.0}/src/cvat_cli/_internal/parsers.py +0 -0
- {cvat_cli-2.68.0 → cvat_cli-2.69.0}/src/cvat_cli/_internal/utils.py +0 -0
- {cvat_cli-2.68.0 → cvat_cli-2.69.0}/src/cvat_cli.egg-info/dependency_links.txt +0 -0
- {cvat_cli-2.68.0 → cvat_cli-2.69.0}/src/cvat_cli.egg-info/entry_points.txt +0 -0
- {cvat_cli-2.68.0 → cvat_cli-2.69.0}/src/cvat_cli.egg-info/top_level.txt +0 -0
|
@@ -1,6 +1,6 @@
|
|
|
1
1
|
Metadata-Version: 2.4
|
|
2
2
|
Name: cvat-cli
|
|
3
|
-
Version: 2.
|
|
3
|
+
Version: 2.69.0
|
|
4
4
|
Summary: Command-line client for CVAT
|
|
5
5
|
Author-email: "CVAT.ai Corporation" <support@cvat.ai>
|
|
6
6
|
License-Expression: MIT
|
|
@@ -9,7 +9,7 @@ Classifier: Programming Language :: Python :: 3
|
|
|
9
9
|
Classifier: Operating System :: OS Independent
|
|
10
10
|
Requires-Python: >=3.10
|
|
11
11
|
Description-Content-Type: text/markdown
|
|
12
|
-
Requires-Dist: cvat-sdk==2.
|
|
12
|
+
Requires-Dist: cvat-sdk==2.69.0
|
|
13
13
|
Requires-Dist: attrs>=24.2.0
|
|
14
14
|
Requires-Dist: Pillow>=10.3.0
|
|
15
15
|
|
|
@@ -14,35 +14,35 @@ import shutil
|
|
|
14
14
|
import tempfile
|
|
15
15
|
import threading
|
|
16
16
|
import time
|
|
17
|
-
from collections import
|
|
18
|
-
from collections.abc import Callable, Generator, Iterator, Sequence
|
|
17
|
+
from collections.abc import Generator, Iterator
|
|
19
18
|
from datetime import datetime, timedelta, timezone
|
|
20
19
|
from http import HTTPStatus
|
|
21
20
|
from pathlib import Path
|
|
22
|
-
from typing import TYPE_CHECKING
|
|
21
|
+
from typing import TYPE_CHECKING
|
|
23
22
|
|
|
24
23
|
import attrs
|
|
25
24
|
import cvat_sdk.auto_annotation as cvataa
|
|
26
25
|
import cvat_sdk.datasets as cvatds
|
|
27
|
-
import PIL.Image
|
|
28
26
|
import urllib3.exceptions
|
|
29
|
-
from cvat_sdk import Client
|
|
30
|
-
from cvat_sdk.auto_annotation.driver import (
|
|
31
|
-
_AnnotationMapper,
|
|
32
|
-
_DetectionFunctionContextImpl,
|
|
33
|
-
_SpecNameMapping,
|
|
34
|
-
)
|
|
27
|
+
from cvat_sdk import Client
|
|
35
28
|
from cvat_sdk.datasets.caching import make_cache_manager
|
|
36
29
|
from cvat_sdk.exceptions import ApiException
|
|
37
30
|
|
|
31
|
+
from .agent_driver import (
|
|
32
|
+
AgentFunctionDriver,
|
|
33
|
+
BadArError,
|
|
34
|
+
IncompatibleFunctionError,
|
|
35
|
+
set_worker_current_function,
|
|
36
|
+
worker_current_function,
|
|
37
|
+
)
|
|
38
|
+
from .agent_driver_detection import AgentDetectionFunctionDriver
|
|
39
|
+
from .agent_driver_tracking import AgentTrackingFunctionDriver, TrackingStateIdGenerator
|
|
38
40
|
from .common import CriticalError, FunctionLoader
|
|
39
41
|
|
|
40
42
|
if TYPE_CHECKING:
|
|
41
43
|
from _typeshed import SupportsReadline
|
|
42
44
|
|
|
43
45
|
FUNCTION_PROVIDER_NATIVE = "native"
|
|
44
|
-
FUNCTION_KIND_DETECTOR = "detector"
|
|
45
|
-
FUNCTION_KIND_TRACKER = "tracker"
|
|
46
46
|
REQUEST_CATEGORY_BATCH = "batch"
|
|
47
47
|
REQUEST_CATEGORY_INTERACTIVE = "interactive"
|
|
48
48
|
|
|
@@ -55,8 +55,6 @@ _DEFAULT_RETRY_DELAY = timedelta(seconds=5)
|
|
|
55
55
|
|
|
56
56
|
_UPDATE_INTERVAL = timedelta(seconds=30)
|
|
57
57
|
|
|
58
|
-
_MAX_AGE_OF_TRACKING_STATE = timedelta(hours=8)
|
|
59
|
-
|
|
60
58
|
|
|
61
59
|
class _ExponentialBackoff:
|
|
62
60
|
def __init__(self, max_delay: timedelta, current_delay: timedelta) -> None:
|
|
@@ -72,7 +70,7 @@ class _ExponentialBackoff:
|
|
|
72
70
|
return delay
|
|
73
71
|
|
|
74
72
|
|
|
75
|
-
class
|
|
73
|
+
class RecoverableExecutor:
|
|
76
74
|
# A wrapper around ProcessPoolExecutor that recreates the underlying
|
|
77
75
|
# executor when a worker crashes.
|
|
78
76
|
def __init__(self, initializer, initargs):
|
|
@@ -107,146 +105,23 @@ class _RecoverableExecutor:
|
|
|
107
105
|
raise
|
|
108
106
|
|
|
109
107
|
|
|
110
|
-
_TrackingStateIdGenerator: TypeAlias = Callable[[], str]
|
|
111
|
-
|
|
112
|
-
|
|
113
108
|
def _default_tracking_state_id_generator() -> str:
|
|
114
109
|
# This is defined as a separate function so that tests can monkeypatch it
|
|
115
110
|
# in order to get deterministic state IDs.
|
|
116
111
|
return secrets.token_urlsafe(32)
|
|
117
112
|
|
|
118
113
|
|
|
119
|
-
|
|
120
|
-
|
|
121
|
-
|
|
122
|
-
|
|
123
|
-
|
|
124
|
-
@attrs.define
|
|
125
|
-
class _ExtendedTrackingState:
|
|
126
|
-
inner_state: Any # the state produced by the AA function
|
|
127
|
-
original_shape_type: str
|
|
128
|
-
original_task_id: int
|
|
129
|
-
original_image_dims: tuple[int, int]
|
|
130
|
-
last_accessed_at: datetime = attrs.field(factory=lambda: datetime.now(tz=timezone.utc))
|
|
131
|
-
|
|
132
|
-
|
|
133
|
-
class _TrackingStateContainer:
|
|
134
|
-
def __init__(self):
|
|
135
|
-
self._id_to_ext_state: OrderedDict[str, _ExtendedTrackingState] = OrderedDict()
|
|
136
|
-
|
|
137
|
-
def store(self, state: Any, shape_type: str, task_id: int, image_dims: tuple[int, int]) -> str:
|
|
138
|
-
state_id = _tracking_state_id_generator()
|
|
139
|
-
self._id_to_ext_state[state_id] = _ExtendedTrackingState(
|
|
140
|
-
inner_state=state,
|
|
141
|
-
original_shape_type=shape_type,
|
|
142
|
-
original_task_id=task_id,
|
|
143
|
-
original_image_dims=image_dims,
|
|
144
|
-
)
|
|
145
|
-
return state_id
|
|
146
|
-
|
|
147
|
-
def retrieve(self, state_id: str, task_id: int, image_dims: tuple[int, int]) -> Any:
|
|
148
|
-
ext_state = self._id_to_ext_state.get(state_id)
|
|
149
|
-
|
|
150
|
-
if not ext_state:
|
|
151
|
-
raise _BadArError(f"Tracking state {state_id!r} not found - possibly expired")
|
|
152
|
-
|
|
153
|
-
if ext_state.original_task_id != task_id:
|
|
154
|
-
# This is a defense-in-depth measure. State IDs are supposed to be unguessable,
|
|
155
|
-
# but even if an attacker manages to obtain one, they will not be able to use it
|
|
156
|
-
# to get any information about a task they don't have access to.
|
|
157
|
-
raise _BadArError(f"Tracking state {state_id!r} is not for task #{task_id}")
|
|
158
|
-
|
|
159
|
-
if image_dims != ext_state.original_image_dims:
|
|
160
|
-
raise _BadArError("Image sizes of the start frame and the current frame are different")
|
|
161
|
-
|
|
162
|
-
ext_state.last_accessed_at = datetime.now(tz=timezone.utc)
|
|
163
|
-
self._id_to_ext_state.move_to_end(state_id)
|
|
164
|
-
|
|
165
|
-
return ext_state.inner_state, ext_state.original_shape_type
|
|
166
|
-
|
|
167
|
-
def prune(self) -> None:
|
|
168
|
-
cutoff = datetime.now(tz=timezone.utc) - _MAX_AGE_OF_TRACKING_STATE
|
|
169
|
-
|
|
170
|
-
while (
|
|
171
|
-
self._id_to_ext_state
|
|
172
|
-
and next(iter(self._id_to_ext_state.values())).last_accessed_at < cutoff
|
|
173
|
-
):
|
|
174
|
-
self._id_to_ext_state.popitem(last=False)
|
|
175
|
-
|
|
176
|
-
|
|
177
|
-
def _worker_init(function_loader: FunctionLoader, state_id_generator):
|
|
178
|
-
global _current_function
|
|
179
|
-
_current_function = function_loader.load()
|
|
180
|
-
|
|
181
|
-
if isinstance(_current_function.spec, cvataa.TrackingFunctionSpec):
|
|
182
|
-
global _tracking_states
|
|
183
|
-
_tracking_states = _TrackingStateContainer()
|
|
114
|
+
def _worker_init(
|
|
115
|
+
function_loader: FunctionLoader, state_id_generator: TrackingStateIdGenerator
|
|
116
|
+
) -> None:
|
|
117
|
+
current_function = function_loader.load()
|
|
118
|
+
set_worker_current_function(current_function)
|
|
184
119
|
|
|
185
|
-
|
|
186
|
-
_tracking_state_id_generator = state_id_generator
|
|
120
|
+
get_function_driver_class(current_function.spec).init_worker(state_id_generator)
|
|
187
121
|
|
|
188
122
|
|
|
189
123
|
def _worker_job_get_function_spec():
|
|
190
|
-
return
|
|
191
|
-
|
|
192
|
-
|
|
193
|
-
def _worker_job_detect(
|
|
194
|
-
context: _DetectionFunctionContextImpl, image: PIL.Image.Image
|
|
195
|
-
) -> list[cvataa.DetectionAnnotation]:
|
|
196
|
-
return _current_function.detect(context, image)
|
|
197
|
-
|
|
198
|
-
|
|
199
|
-
def _worker_job_init_tracking(
|
|
200
|
-
task_id: int,
|
|
201
|
-
image: PIL.Image.Image,
|
|
202
|
-
shapes: list[cvataa.TrackableShape],
|
|
203
|
-
) -> list[str]:
|
|
204
|
-
_tracking_states.prune()
|
|
205
|
-
|
|
206
|
-
if hasattr(_current_function, "preprocess_image"):
|
|
207
|
-
pp_image = _current_function.preprocess_image(_TrackingFunctionContextImpl(), image)
|
|
208
|
-
else:
|
|
209
|
-
pp_image = image
|
|
210
|
-
|
|
211
|
-
return [
|
|
212
|
-
_tracking_states.store(
|
|
213
|
-
state=_current_function.init_tracking_state(
|
|
214
|
-
_TrackingFunctionShapeContextImpl(original_shape_type=shape.type), pp_image, shape
|
|
215
|
-
),
|
|
216
|
-
shape_type=shape.type,
|
|
217
|
-
task_id=task_id,
|
|
218
|
-
image_dims=image.size,
|
|
219
|
-
)
|
|
220
|
-
for shape in shapes
|
|
221
|
-
]
|
|
222
|
-
|
|
223
|
-
|
|
224
|
-
def _worker_job_track(
|
|
225
|
-
task_id: int, image: PIL.Image.Image, states: list[str]
|
|
226
|
-
) -> list[cvataa.TrackableShape | None]:
|
|
227
|
-
_tracking_states.prune()
|
|
228
|
-
|
|
229
|
-
pp_image = _current_function.preprocess_image(_TrackingFunctionContextImpl(), image)
|
|
230
|
-
|
|
231
|
-
def track(state_id):
|
|
232
|
-
inner_state, original_shape_type = _tracking_states.retrieve(
|
|
233
|
-
state_id=state_id, task_id=task_id, image_dims=image.size
|
|
234
|
-
)
|
|
235
|
-
|
|
236
|
-
output_shape = _current_function.track(
|
|
237
|
-
_TrackingFunctionShapeContextImpl(original_shape_type=original_shape_type),
|
|
238
|
-
pp_image,
|
|
239
|
-
inner_state,
|
|
240
|
-
)
|
|
241
|
-
|
|
242
|
-
if output_shape and output_shape.type != original_shape_type:
|
|
243
|
-
raise cvataa.BadFunctionError(
|
|
244
|
-
f"function output shape of type {output_shape.type!r}, "
|
|
245
|
-
f"but original shape was of type {original_shape_type!r}"
|
|
246
|
-
)
|
|
247
|
-
return output_shape
|
|
248
|
-
|
|
249
|
-
return list(map(track, states))
|
|
124
|
+
return worker_current_function().spec
|
|
250
125
|
|
|
251
126
|
|
|
252
127
|
@attrs.frozen
|
|
@@ -348,33 +223,25 @@ def _parse_event_stream(
|
|
|
348
223
|
yield _NewReconnectionDelay(timedelta(milliseconds=int(field_value)))
|
|
349
224
|
|
|
350
225
|
|
|
351
|
-
|
|
352
|
-
|
|
353
|
-
|
|
226
|
+
def get_function_driver_class(function_spec: object) -> type[AgentFunctionDriver]:
|
|
227
|
+
if isinstance(function_spec, cvataa.DetectionFunctionSpec):
|
|
228
|
+
return AgentDetectionFunctionDriver
|
|
354
229
|
|
|
355
|
-
|
|
356
|
-
|
|
357
|
-
pass
|
|
230
|
+
if isinstance(function_spec, cvataa.TrackingFunctionSpec):
|
|
231
|
+
return AgentTrackingFunctionDriver
|
|
358
232
|
|
|
359
|
-
|
|
360
|
-
class _TrackingFunctionContextImpl(cvataa.TrackingFunctionContext):
|
|
361
|
-
pass
|
|
362
|
-
|
|
363
|
-
|
|
364
|
-
@attrs.frozen(kw_only=True)
|
|
365
|
-
class _TrackingFunctionShapeContextImpl(cvataa.TrackingFunctionShapeContext):
|
|
366
|
-
original_shape_type: str
|
|
233
|
+
raise CriticalError(f"Unsupported function spec type: {type(function_spec).__name__}")
|
|
367
234
|
|
|
368
235
|
|
|
369
236
|
class _Agent:
|
|
370
|
-
def __init__(self, client: Client, executor:
|
|
237
|
+
def __init__(self, client: Client, executor: RecoverableExecutor, function_id: int):
|
|
371
238
|
self._rng = random.Random() # nosec
|
|
372
239
|
|
|
373
240
|
self._client = client
|
|
374
|
-
self._executor = executor
|
|
375
241
|
self._function_id = function_id
|
|
376
|
-
|
|
377
|
-
|
|
242
|
+
function_spec = executor.result(executor.submit(_worker_job_get_function_spec))
|
|
243
|
+
self._function_driver = get_function_driver_class(function_spec)(
|
|
244
|
+
client, executor, function_spec
|
|
378
245
|
)
|
|
379
246
|
|
|
380
247
|
_, response = self._client.api_client.call_api(
|
|
@@ -422,91 +289,18 @@ class _Agent:
|
|
|
422
289
|
)
|
|
423
290
|
|
|
424
291
|
try:
|
|
425
|
-
if
|
|
426
|
-
|
|
427
|
-
|
|
428
|
-
|
|
429
|
-
self._validate_tracking_function_compatibility(remote_function)
|
|
430
|
-
self._calculate_result_for_ar = self._calculate_result_for_tracking_ar
|
|
431
|
-
else:
|
|
432
|
-
raise CriticalError(
|
|
433
|
-
f"Unsupported function spec type: {type(self._function_spec).__name__}"
|
|
292
|
+
if remote_function["kind"] != self._function_driver.FUNCTION_KIND:
|
|
293
|
+
raise IncompatibleFunctionError(
|
|
294
|
+
f"kind is {remote_function['kind']!r} "
|
|
295
|
+
f"(expected {self._function_driver.FUNCTION_KIND!r})."
|
|
434
296
|
)
|
|
435
|
-
|
|
297
|
+
|
|
298
|
+
self._function_driver.validate_function_compatibility(remote_function)
|
|
299
|
+
except IncompatibleFunctionError as ex:
|
|
436
300
|
raise CriticalError(
|
|
437
301
|
f"Function #{function_id} is incompatible with function object: {ex}"
|
|
438
302
|
) from ex
|
|
439
303
|
|
|
440
|
-
def _validate_detection_function_compatibility(self, remote_function: dict) -> None:
|
|
441
|
-
self._validate_remote_function_kind(remote_function, FUNCTION_KIND_DETECTOR)
|
|
442
|
-
|
|
443
|
-
labels_by_name = {label.name: label for label in self._function_spec.labels}
|
|
444
|
-
|
|
445
|
-
for remote_label in remote_function["labels_v2"]:
|
|
446
|
-
label_desc = f"label {remote_label['name']!r}"
|
|
447
|
-
label = labels_by_name.get(remote_label["name"])
|
|
448
|
-
|
|
449
|
-
self._validate_sublabel_compatibility(remote_label, label, label_desc)
|
|
450
|
-
|
|
451
|
-
sublabels_by_name = {sl.name: sl for sl in getattr(label, "sublabels", [])}
|
|
452
|
-
|
|
453
|
-
for remote_sl in remote_label.get("sublabels", []):
|
|
454
|
-
sl_desc = f"sublabel {remote_sl['name']!r} of {label_desc}"
|
|
455
|
-
sl = sublabels_by_name.get(remote_sl["name"])
|
|
456
|
-
|
|
457
|
-
self._validate_sublabel_compatibility(remote_sl, sl, sl_desc)
|
|
458
|
-
|
|
459
|
-
def _validate_sublabel_compatibility(
|
|
460
|
-
self, remote_sl: dict, sl: models.Sublabel | None, sl_desc: str
|
|
461
|
-
):
|
|
462
|
-
if not sl:
|
|
463
|
-
raise _IncompatibleFunctionError(f"{sl_desc} is not supported.")
|
|
464
|
-
|
|
465
|
-
if remote_sl["type"] not in {"any", "unknown"} and remote_sl["type"] != sl.type:
|
|
466
|
-
raise _IncompatibleFunctionError(
|
|
467
|
-
f"{sl_desc} has type {remote_sl['type']!r}, "
|
|
468
|
-
f"but the function object declares type {sl.type!r}."
|
|
469
|
-
)
|
|
470
|
-
|
|
471
|
-
attrs_by_name = {attr.name: attr for attr in getattr(sl, "attributes", [])}
|
|
472
|
-
|
|
473
|
-
for remote_attr in remote_sl["attributes"]:
|
|
474
|
-
attr_desc = f"attribute {remote_attr['name']!r} of {sl_desc}"
|
|
475
|
-
attr = attrs_by_name.get(remote_attr["name"])
|
|
476
|
-
|
|
477
|
-
if not attr:
|
|
478
|
-
raise _IncompatibleFunctionError(f"{attr_desc} is not supported.")
|
|
479
|
-
|
|
480
|
-
if remote_attr["input_type"] != attr.input_type.value:
|
|
481
|
-
raise _IncompatibleFunctionError(
|
|
482
|
-
f"{attr_desc} has input type {remote_attr['input_type']!r},"
|
|
483
|
-
f" but the function object declares input type {attr.input_type.value!r}."
|
|
484
|
-
)
|
|
485
|
-
|
|
486
|
-
if remote_attr["values"] != attr.values:
|
|
487
|
-
raise _IncompatibleFunctionError(
|
|
488
|
-
f"{attr_desc} has values {remote_attr['values']!r},"
|
|
489
|
-
f" but the function object declares values {attr.values!r}."
|
|
490
|
-
)
|
|
491
|
-
|
|
492
|
-
def _validate_tracking_function_compatibility(self, remote_function: dict) -> None:
|
|
493
|
-
self._validate_remote_function_kind(remote_function, FUNCTION_KIND_TRACKER)
|
|
494
|
-
|
|
495
|
-
remote_supported_shape_types = frozenset(remote_function["supported_shape_types"])
|
|
496
|
-
unsupported = remote_supported_shape_types - self._function_spec.supported_shape_types
|
|
497
|
-
|
|
498
|
-
if unsupported:
|
|
499
|
-
raise _IncompatibleFunctionError(
|
|
500
|
-
"the function object does not support the following shape types: "
|
|
501
|
-
+ ", ".join(map(repr, unsupported))
|
|
502
|
-
)
|
|
503
|
-
|
|
504
|
-
def _validate_remote_function_kind(self, remote_function: dict, expected_kind: str) -> None:
|
|
505
|
-
if remote_function["kind"] != expected_kind:
|
|
506
|
-
raise _IncompatibleFunctionError(
|
|
507
|
-
f"kind is {remote_function['kind']!r} (expected {expected_kind!r})."
|
|
508
|
-
)
|
|
509
|
-
|
|
510
304
|
def _wait_between_polls(self):
|
|
511
305
|
# offset the interval randomly to avoid synchronization between workers
|
|
512
306
|
timeout_multiplier = self._rng.uniform(1 - _JITTER_AMOUNT, 1 + _JITTER_AMOUNT)
|
|
@@ -678,9 +472,23 @@ class _Agent:
|
|
|
678
472
|
)
|
|
679
473
|
self._client.logger.debug("AR %r parameters: %r", ar_id, ar_params)
|
|
680
474
|
|
|
681
|
-
|
|
682
|
-
result = self._calculate_result_for_ar(ar_id, ar_params)
|
|
475
|
+
last_update_timestamp = datetime.now(tz=timezone.utc)
|
|
683
476
|
|
|
477
|
+
def check_in(*, current_progress: float) -> None:
|
|
478
|
+
nonlocal last_update_timestamp
|
|
479
|
+
current_timestamp = datetime.now(tz=timezone.utc)
|
|
480
|
+
|
|
481
|
+
if current_timestamp >= last_update_timestamp + _UPDATE_INTERVAL:
|
|
482
|
+
self._update_ar(ar_id, current_progress)
|
|
483
|
+
last_update_timestamp = current_timestamp
|
|
484
|
+
|
|
485
|
+
# Interactive requests are time sensitive, so if there are any,
|
|
486
|
+
# we have to put the current AR on hold and process them ASAP.
|
|
487
|
+
self._process_available_ars(REQUEST_CATEGORY_INTERACTIVE)
|
|
488
|
+
|
|
489
|
+
try:
|
|
490
|
+
with self._task_cache_limiter.using_cache_for_task(ar_params["task"]):
|
|
491
|
+
result = self._function_driver.calculate_result_for_ar(ar_params, check_in)
|
|
684
492
|
self._complete_ar(ar_id, result)
|
|
685
493
|
except Exception as ex:
|
|
686
494
|
self._client.logger.error("Failed to process AR %r", ar_id, exc_info=True)
|
|
@@ -706,7 +514,7 @@ class _Agent:
|
|
|
706
514
|
error_message = "Failed to make an HTTP request"
|
|
707
515
|
elif isinstance(ex, cvataa.BadFunctionError):
|
|
708
516
|
error_message = "Underlying function returned incorrect result: " + str(ex)
|
|
709
|
-
elif isinstance(ex,
|
|
517
|
+
elif isinstance(ex, BadArError):
|
|
710
518
|
error_message = "Invalid annotation request: " + str(ex)
|
|
711
519
|
elif isinstance(ex, concurrent.futures.BrokenExecutor):
|
|
712
520
|
error_message = "Worker process crashed"
|
|
@@ -784,164 +592,6 @@ class _Agent:
|
|
|
784
592
|
response_data = json.loads(response.data)
|
|
785
593
|
return response_data["ar_assignment"]
|
|
786
594
|
|
|
787
|
-
def _calculate_result_for_detection_ar(self, ar_id: str, ar_params) -> dict[str, Any]:
|
|
788
|
-
if ar_params["type"] == "annotate_task":
|
|
789
|
-
with self._task_cache_limiter.using_cache_for_task(ar_params["task"]):
|
|
790
|
-
return self._calculate_result_for_annotate_task_ar(ar_id, ar_params)
|
|
791
|
-
elif ar_params["type"] == "annotate_frame":
|
|
792
|
-
with self._task_cache_limiter.using_cache_for_task(ar_params["task"]):
|
|
793
|
-
return self._calculate_result_for_annotate_frame_ar(ar_id, ar_params)
|
|
794
|
-
else:
|
|
795
|
-
raise _BadArError(f"unsupported type: {ar_params['type']!r}")
|
|
796
|
-
|
|
797
|
-
def _create_annotation_mapper_for_detection_ar(
|
|
798
|
-
self, ar_params: dict, ds_labels: Sequence[models.ILabel]
|
|
799
|
-
) -> _AnnotationMapper:
|
|
800
|
-
spec_nm = _SpecNameMapping.from_api(
|
|
801
|
-
{
|
|
802
|
-
k: models.LabelMappingEntryRequest._from_openapi_data(**v)
|
|
803
|
-
for k, v in ar_params["mapping"].items()
|
|
804
|
-
}
|
|
805
|
-
)
|
|
806
|
-
|
|
807
|
-
return _AnnotationMapper(
|
|
808
|
-
self._client.logger,
|
|
809
|
-
self._function_spec.labels,
|
|
810
|
-
ds_labels,
|
|
811
|
-
allow_unmatched_labels=False,
|
|
812
|
-
spec_nm=spec_nm,
|
|
813
|
-
conv_mask_to_poly=ar_params["conv_mask_to_poly"],
|
|
814
|
-
)
|
|
815
|
-
|
|
816
|
-
def _create_detection_function_context(
|
|
817
|
-
self, ar_params: dict, frame_name: str
|
|
818
|
-
) -> cvataa.DetectionFunctionContext:
|
|
819
|
-
return _DetectionFunctionContextImpl(
|
|
820
|
-
frame_name=frame_name,
|
|
821
|
-
conf_threshold=ar_params["threshold"],
|
|
822
|
-
conv_mask_to_poly=ar_params["conv_mask_to_poly"],
|
|
823
|
-
)
|
|
824
|
-
|
|
825
|
-
def _calculate_result_for_annotate_task_ar(self, ar_id: str, ar_params) -> dict[str, Any]:
|
|
826
|
-
ds = cvatds.TaskDataset(
|
|
827
|
-
self._client,
|
|
828
|
-
ar_params["task"],
|
|
829
|
-
load_annotations=False,
|
|
830
|
-
media_download_policy=cvatds.MediaDownloadPolicy.FETCH_CHUNKS_ON_DEMAND,
|
|
831
|
-
)
|
|
832
|
-
|
|
833
|
-
# Fetching the dataset might take a while, so do a progress update to let the server
|
|
834
|
-
# know we're still alive.
|
|
835
|
-
self._update_ar(ar_id, 0)
|
|
836
|
-
last_update_timestamp = datetime.now(tz=timezone.utc)
|
|
837
|
-
|
|
838
|
-
mapper = self._create_annotation_mapper_for_detection_ar(ar_params, ds.labels)
|
|
839
|
-
|
|
840
|
-
all_annotations = models.PatchedLabeledDataRequest(tags=[], shapes=[])
|
|
841
|
-
|
|
842
|
-
with ds.iter_samples(temporary_chunks=True) as samples:
|
|
843
|
-
for sample_index, sample in enumerate(samples):
|
|
844
|
-
context = self._create_detection_function_context(ar_params, sample.frame_name)
|
|
845
|
-
annotations = self._executor.result(
|
|
846
|
-
self._executor.submit(_worker_job_detect, context, sample.media.load_image())
|
|
847
|
-
)
|
|
848
|
-
|
|
849
|
-
tags, shapes = mapper.validate_and_remap(annotations, sample.frame_index)
|
|
850
|
-
all_annotations.tags.extend(tags)
|
|
851
|
-
all_annotations.shapes.extend(shapes)
|
|
852
|
-
|
|
853
|
-
current_timestamp = datetime.now(tz=timezone.utc)
|
|
854
|
-
|
|
855
|
-
if current_timestamp >= last_update_timestamp + _UPDATE_INTERVAL:
|
|
856
|
-
self._update_ar(ar_id, (sample_index + 1) / len(ds.samples))
|
|
857
|
-
last_update_timestamp = current_timestamp
|
|
858
|
-
|
|
859
|
-
# Interactive requests are time sensitive, so if there are any,
|
|
860
|
-
# we have to put the current AR on hold and process them ASAP.
|
|
861
|
-
self._process_available_ars(REQUEST_CATEGORY_INTERACTIVE)
|
|
862
|
-
|
|
863
|
-
return {"annotations": all_annotations}
|
|
864
|
-
|
|
865
|
-
def _calculate_result_for_annotate_frame_ar(self, ar_id: str, ar_params) -> dict[str, Any]:
|
|
866
|
-
sample, ds_labels = self._get_sample_from_ar_params(ar_params)
|
|
867
|
-
|
|
868
|
-
mapper = self._create_annotation_mapper_for_detection_ar(ar_params, ds_labels)
|
|
869
|
-
|
|
870
|
-
context = self._create_detection_function_context(ar_params, sample.frame_name)
|
|
871
|
-
|
|
872
|
-
annotations = self._executor.result(
|
|
873
|
-
self._executor.submit(_worker_job_detect, context, sample.media.load_image())
|
|
874
|
-
)
|
|
875
|
-
|
|
876
|
-
tags, shapes = mapper.validate_and_remap(annotations, sample.frame_index)
|
|
877
|
-
return {"annotations": models.PatchedLabeledDataRequest(tags=tags, shapes=shapes)}
|
|
878
|
-
|
|
879
|
-
def _calculate_result_for_tracking_ar(self, ar_id: str, ar_params) -> dict[str, Any]:
|
|
880
|
-
if ar_params["type"] == "init_tracking":
|
|
881
|
-
with self._task_cache_limiter.using_cache_for_task(ar_params["task"]):
|
|
882
|
-
return self._calculate_result_for_init_tracking_ar(ar_id, ar_params)
|
|
883
|
-
elif ar_params["type"] == "track":
|
|
884
|
-
with self._task_cache_limiter.using_cache_for_task(ar_params["task"]):
|
|
885
|
-
return self._calculate_result_for_track_ar(ar_id, ar_params)
|
|
886
|
-
else:
|
|
887
|
-
raise _BadArError(f"unsupported type: {ar_params['type']!r}")
|
|
888
|
-
|
|
889
|
-
def _calculate_result_for_init_tracking_ar(self, ar_id: str, ar_params) -> dict[str, Any]:
|
|
890
|
-
sample, _ = self._get_sample_from_ar_params(ar_params)
|
|
891
|
-
|
|
892
|
-
def convert_shape(shape: dict) -> cvataa.TrackableShape:
|
|
893
|
-
if shape["type"] not in self._function_spec.supported_shape_types:
|
|
894
|
-
raise _BadArError(f"Unsupported shape type {shape['type']!r}")
|
|
895
|
-
return cvataa.TrackableShape(type=shape["type"], points=shape["points"])
|
|
896
|
-
|
|
897
|
-
shapes = list(map(convert_shape, ar_params["shapes"]))
|
|
898
|
-
|
|
899
|
-
states = self._executor.result(
|
|
900
|
-
self._executor.submit(
|
|
901
|
-
_worker_job_init_tracking,
|
|
902
|
-
ar_params["task"],
|
|
903
|
-
sample.media.load_image(),
|
|
904
|
-
shapes,
|
|
905
|
-
)
|
|
906
|
-
)
|
|
907
|
-
|
|
908
|
-
return {"states": states}
|
|
909
|
-
|
|
910
|
-
def _calculate_result_for_track_ar(self, ar_id: str, ar_params) -> dict[str, Any]:
|
|
911
|
-
sample, _ = self._get_sample_from_ar_params(ar_params)
|
|
912
|
-
|
|
913
|
-
states = ar_params["states"]
|
|
914
|
-
shapes = self._executor.result(
|
|
915
|
-
self._executor.submit(
|
|
916
|
-
_worker_job_track, ar_params["task"], sample.media.load_image(), states
|
|
917
|
-
)
|
|
918
|
-
)
|
|
919
|
-
|
|
920
|
-
return {
|
|
921
|
-
"states": states,
|
|
922
|
-
"shapes": [attrs.asdict(shape) if shape else None for shape in shapes],
|
|
923
|
-
}
|
|
924
|
-
|
|
925
|
-
def _get_sample_from_ar_params(self, ar_params):
|
|
926
|
-
ds = cvatds.TaskDataset(
|
|
927
|
-
self._client,
|
|
928
|
-
ar_params["task"],
|
|
929
|
-
load_annotations=False,
|
|
930
|
-
media_download_policy=cvatds.MediaDownloadPolicy.FETCH_FRAMES_ON_DEMAND,
|
|
931
|
-
)
|
|
932
|
-
|
|
933
|
-
frame_index = ar_params["frame"]
|
|
934
|
-
|
|
935
|
-
# Since ds.samples excludes deleted frames, we can't just do sample = ds.samples[frame_index].
|
|
936
|
-
# Once we drop Python 3.9, we can change this to use bisect instead of the linear search.
|
|
937
|
-
for sample in ds.samples:
|
|
938
|
-
if sample.frame_index == frame_index:
|
|
939
|
-
break
|
|
940
|
-
else:
|
|
941
|
-
raise _BadArError(f"Frame with index {frame_index} does not exist in the task")
|
|
942
|
-
|
|
943
|
-
return sample, ds.labels
|
|
944
|
-
|
|
945
595
|
def _update_ar(self, ar_id: str, progress: float) -> None:
|
|
946
596
|
self._client.logger.info("Updating AR %r progress to %.2f%%...", ar_id, progress * 100)
|
|
947
597
|
|
|
@@ -992,7 +642,7 @@ def run_agent(
|
|
|
992
642
|
client: Client, function_loader: FunctionLoader, function_id: int, *, burst: bool
|
|
993
643
|
) -> None:
|
|
994
644
|
with (
|
|
995
|
-
|
|
645
|
+
RecoverableExecutor(
|
|
996
646
|
initializer=_worker_init,
|
|
997
647
|
initargs=[function_loader, _default_tracking_state_id_generator],
|
|
998
648
|
) as executor,
|
|
@@ -0,0 +1,95 @@
|
|
|
1
|
+
# Copyright (C) CVAT.ai Corporation
|
|
2
|
+
#
|
|
3
|
+
# SPDX-License-Identifier: MIT
|
|
4
|
+
|
|
5
|
+
from __future__ import annotations
|
|
6
|
+
|
|
7
|
+
from collections.abc import Callable
|
|
8
|
+
from typing import TYPE_CHECKING, Any, ClassVar, Protocol
|
|
9
|
+
|
|
10
|
+
import cvat_sdk.auto_annotation as cvataa
|
|
11
|
+
import cvat_sdk.datasets as cvatds
|
|
12
|
+
from cvat_sdk import Client
|
|
13
|
+
from typing_extensions import Self
|
|
14
|
+
|
|
15
|
+
if TYPE_CHECKING:
|
|
16
|
+
from .agent import RecoverableExecutor
|
|
17
|
+
from .agent_driver_tracking import TrackingStateIdGenerator
|
|
18
|
+
|
|
19
|
+
|
|
20
|
+
class BadArError(Exception):
|
|
21
|
+
pass
|
|
22
|
+
|
|
23
|
+
|
|
24
|
+
class IncompatibleFunctionError(Exception):
|
|
25
|
+
# This should only be thrown from inside validate_function_compatibility methods.
|
|
26
|
+
pass
|
|
27
|
+
|
|
28
|
+
|
|
29
|
+
_current_function: cvataa.AutoAnnotationFunction
|
|
30
|
+
|
|
31
|
+
|
|
32
|
+
def worker_current_function() -> cvataa.AutoAnnotationFunction:
|
|
33
|
+
return _current_function
|
|
34
|
+
|
|
35
|
+
|
|
36
|
+
def set_worker_current_function(func: cvataa.AutoAnnotationFunction) -> None:
|
|
37
|
+
global _current_function
|
|
38
|
+
_current_function = func
|
|
39
|
+
|
|
40
|
+
|
|
41
|
+
class CheckInCallback(Protocol):
|
|
42
|
+
def __call__(self, *, current_progress: float) -> None:
|
|
43
|
+
"""
|
|
44
|
+
Temporarily suspends the processing of the current AR to report progress back to the server
|
|
45
|
+
and optionally process any pending interactive ARs. Should only be called during batch AR
|
|
46
|
+
processing (interactive ARs should be processed quickly enough to not require this).
|
|
47
|
+
"""
|
|
48
|
+
|
|
49
|
+
|
|
50
|
+
class AgentFunctionDriver:
|
|
51
|
+
FUNCTION_KIND: ClassVar[str]
|
|
52
|
+
|
|
53
|
+
def __init__(self, client: Client, executor: RecoverableExecutor, function_spec: object):
|
|
54
|
+
self._client = client
|
|
55
|
+
self._executor = executor
|
|
56
|
+
self._function_spec = function_spec
|
|
57
|
+
|
|
58
|
+
@classmethod
|
|
59
|
+
def init_worker(cls, state_id_generator: TrackingStateIdGenerator) -> None:
|
|
60
|
+
pass
|
|
61
|
+
|
|
62
|
+
def validate_function_compatibility(self, remote_function: dict) -> None:
|
|
63
|
+
raise NotImplementedError
|
|
64
|
+
|
|
65
|
+
_CALCULATE_RESULT_PER_AR_TYPE: ClassVar[
|
|
66
|
+
dict[str, Callable[[Self, dict[str, Any], CheckInCallback], dict[str, Any]]]
|
|
67
|
+
]
|
|
68
|
+
|
|
69
|
+
def calculate_result_for_ar(
|
|
70
|
+
self, ar_params: dict[str, Any], check_in: CheckInCallback
|
|
71
|
+
) -> dict[str, Any]:
|
|
72
|
+
if calc := self._CALCULATE_RESULT_PER_AR_TYPE.get(ar_params["type"]):
|
|
73
|
+
return calc(self, ar_params, check_in)
|
|
74
|
+
|
|
75
|
+
raise BadArError(f"unsupported type: {ar_params['type']!r}")
|
|
76
|
+
|
|
77
|
+
def _get_sample_from_ar_params(self, ar_params):
|
|
78
|
+
ds = cvatds.TaskDataset(
|
|
79
|
+
self._client,
|
|
80
|
+
ar_params["task"],
|
|
81
|
+
load_annotations=False,
|
|
82
|
+
media_download_policy=cvatds.MediaDownloadPolicy.FETCH_FRAMES_ON_DEMAND,
|
|
83
|
+
)
|
|
84
|
+
|
|
85
|
+
frame_index = ar_params["frame"]
|
|
86
|
+
|
|
87
|
+
# Since ds.samples excludes deleted frames, we can't just do sample = ds.samples[frame_index].
|
|
88
|
+
# Once we drop Python 3.9, we can change this to use bisect instead of the linear search.
|
|
89
|
+
for sample in ds.samples:
|
|
90
|
+
if sample.frame_index == frame_index:
|
|
91
|
+
break
|
|
92
|
+
else:
|
|
93
|
+
raise BadArError(f"Frame with index {frame_index} does not exist in the task")
|
|
94
|
+
|
|
95
|
+
return sample, ds.labels
|
|
@@ -0,0 +1,167 @@
|
|
|
1
|
+
# Copyright (C) CVAT.ai Corporation
|
|
2
|
+
#
|
|
3
|
+
# SPDX-License-Identifier: MIT
|
|
4
|
+
|
|
5
|
+
from collections.abc import Sequence
|
|
6
|
+
from typing import Any, cast
|
|
7
|
+
|
|
8
|
+
import cvat_sdk.auto_annotation as cvataa
|
|
9
|
+
import cvat_sdk.datasets as cvatds
|
|
10
|
+
import PIL.Image
|
|
11
|
+
from cvat_sdk import models
|
|
12
|
+
from cvat_sdk.auto_annotation.driver import (
|
|
13
|
+
_AnnotationMapper,
|
|
14
|
+
_DetectionFunctionContextImpl,
|
|
15
|
+
_SpecNameMapping,
|
|
16
|
+
)
|
|
17
|
+
|
|
18
|
+
from .agent_driver import (
|
|
19
|
+
AgentFunctionDriver,
|
|
20
|
+
CheckInCallback,
|
|
21
|
+
IncompatibleFunctionError,
|
|
22
|
+
worker_current_function,
|
|
23
|
+
)
|
|
24
|
+
|
|
25
|
+
|
|
26
|
+
def _worker_job_detect(
|
|
27
|
+
context: _DetectionFunctionContextImpl, image: PIL.Image.Image
|
|
28
|
+
) -> Sequence[cvataa.DetectionAnnotation]:
|
|
29
|
+
current_function = cast(cvataa.DetectionFunction, worker_current_function())
|
|
30
|
+
return current_function.detect(context, image)
|
|
31
|
+
|
|
32
|
+
|
|
33
|
+
class AgentDetectionFunctionDriver(AgentFunctionDriver):
|
|
34
|
+
FUNCTION_KIND = "detector"
|
|
35
|
+
_function_spec: cvataa.DetectionFunctionSpec
|
|
36
|
+
|
|
37
|
+
def _validate_sublabel_compatibility(
|
|
38
|
+
self, remote_sl: dict, sl: models.Sublabel | None, sl_desc: str
|
|
39
|
+
):
|
|
40
|
+
if not sl:
|
|
41
|
+
raise IncompatibleFunctionError(f"{sl_desc} is not supported.")
|
|
42
|
+
|
|
43
|
+
if remote_sl["type"] not in {"any", "unknown"} and remote_sl["type"] != sl.type:
|
|
44
|
+
raise IncompatibleFunctionError(
|
|
45
|
+
f"{sl_desc} has type {remote_sl['type']!r}, "
|
|
46
|
+
f"but the function object declares type {sl.type!r}."
|
|
47
|
+
)
|
|
48
|
+
|
|
49
|
+
attrs_by_name = {attr.name: attr for attr in getattr(sl, "attributes", [])}
|
|
50
|
+
|
|
51
|
+
for remote_attr in remote_sl["attributes"]:
|
|
52
|
+
attr_desc = f"attribute {remote_attr['name']!r} of {sl_desc}"
|
|
53
|
+
attr = attrs_by_name.get(remote_attr["name"])
|
|
54
|
+
|
|
55
|
+
if not attr:
|
|
56
|
+
raise IncompatibleFunctionError(f"{attr_desc} is not supported.")
|
|
57
|
+
|
|
58
|
+
if remote_attr["input_type"] != attr.input_type.value:
|
|
59
|
+
raise IncompatibleFunctionError(
|
|
60
|
+
f"{attr_desc} has input type {remote_attr['input_type']!r},"
|
|
61
|
+
f" but the function object declares input type {attr.input_type.value!r}."
|
|
62
|
+
)
|
|
63
|
+
|
|
64
|
+
if remote_attr["values"] != attr.values:
|
|
65
|
+
raise IncompatibleFunctionError(
|
|
66
|
+
f"{attr_desc} has values {remote_attr['values']!r},"
|
|
67
|
+
f" but the function object declares values {attr.values!r}."
|
|
68
|
+
)
|
|
69
|
+
|
|
70
|
+
def validate_function_compatibility(self, remote_function: dict) -> None:
|
|
71
|
+
labels_by_name = {label.name: label for label in self._function_spec.labels}
|
|
72
|
+
|
|
73
|
+
for remote_label in remote_function["labels_v2"]:
|
|
74
|
+
label_desc = f"label {remote_label['name']!r}"
|
|
75
|
+
label = labels_by_name.get(remote_label["name"])
|
|
76
|
+
|
|
77
|
+
self._validate_sublabel_compatibility(remote_label, label, label_desc)
|
|
78
|
+
|
|
79
|
+
sublabels_by_name = {sl.name: sl for sl in getattr(label, "sublabels", [])}
|
|
80
|
+
|
|
81
|
+
for remote_sl in remote_label.get("sublabels", []):
|
|
82
|
+
sl_desc = f"sublabel {remote_sl['name']!r} of {label_desc}"
|
|
83
|
+
sl = sublabels_by_name.get(remote_sl["name"])
|
|
84
|
+
|
|
85
|
+
self._validate_sublabel_compatibility(remote_sl, sl, sl_desc)
|
|
86
|
+
|
|
87
|
+
def _create_annotation_mapper_for_detection_ar(
|
|
88
|
+
self, ar_params: dict, ds_labels: Sequence[models.ILabel]
|
|
89
|
+
) -> _AnnotationMapper:
|
|
90
|
+
spec_nm = _SpecNameMapping.from_api(
|
|
91
|
+
{
|
|
92
|
+
k: models.LabelMappingEntryRequest._from_openapi_data(**v)
|
|
93
|
+
for k, v in ar_params["mapping"].items()
|
|
94
|
+
}
|
|
95
|
+
)
|
|
96
|
+
|
|
97
|
+
return _AnnotationMapper(
|
|
98
|
+
self._client.logger,
|
|
99
|
+
self._function_spec.labels,
|
|
100
|
+
ds_labels,
|
|
101
|
+
allow_unmatched_labels=False,
|
|
102
|
+
spec_nm=spec_nm,
|
|
103
|
+
conv_mask_to_poly=ar_params["conv_mask_to_poly"],
|
|
104
|
+
)
|
|
105
|
+
|
|
106
|
+
def _create_detection_function_context(
|
|
107
|
+
self, ar_params: dict, frame_name: str
|
|
108
|
+
) -> cvataa.DetectionFunctionContext:
|
|
109
|
+
return _DetectionFunctionContextImpl(
|
|
110
|
+
frame_name=frame_name,
|
|
111
|
+
conf_threshold=ar_params["threshold"],
|
|
112
|
+
conv_mask_to_poly=ar_params["conv_mask_to_poly"],
|
|
113
|
+
)
|
|
114
|
+
|
|
115
|
+
def _calculate_result_for_annotate_task_ar(
|
|
116
|
+
self, ar_params: dict[str, Any], check_in: CheckInCallback
|
|
117
|
+
) -> dict[str, Any]:
|
|
118
|
+
ds = cvatds.TaskDataset(
|
|
119
|
+
self._client,
|
|
120
|
+
ar_params["task"],
|
|
121
|
+
load_annotations=False,
|
|
122
|
+
media_download_policy=cvatds.MediaDownloadPolicy.FETCH_CHUNKS_ON_DEMAND,
|
|
123
|
+
)
|
|
124
|
+
|
|
125
|
+
# Fetching the dataset might take a while, so check in to let the server
|
|
126
|
+
# know we're still alive.
|
|
127
|
+
check_in(current_progress=0)
|
|
128
|
+
|
|
129
|
+
mapper = self._create_annotation_mapper_for_detection_ar(ar_params, ds.labels)
|
|
130
|
+
|
|
131
|
+
all_annotations = models.PatchedLabeledDataRequest(tags=[], shapes=[])
|
|
132
|
+
|
|
133
|
+
with ds.iter_samples(temporary_chunks=True) as samples:
|
|
134
|
+
for sample_index, sample in enumerate(samples):
|
|
135
|
+
context = self._create_detection_function_context(ar_params, sample.frame_name)
|
|
136
|
+
annotations = self._executor.result(
|
|
137
|
+
self._executor.submit(_worker_job_detect, context, sample.media.load_image())
|
|
138
|
+
)
|
|
139
|
+
|
|
140
|
+
tags, shapes = mapper.validate_and_remap(annotations, sample.frame_index)
|
|
141
|
+
all_annotations.tags.extend(tags)
|
|
142
|
+
all_annotations.shapes.extend(shapes)
|
|
143
|
+
|
|
144
|
+
check_in(current_progress=(sample_index + 1) / len(ds.samples))
|
|
145
|
+
|
|
146
|
+
return {"annotations": all_annotations}
|
|
147
|
+
|
|
148
|
+
def _calculate_result_for_annotate_frame_ar(
|
|
149
|
+
self, ar_params: dict[str, Any], check_in: object
|
|
150
|
+
) -> dict[str, Any]:
|
|
151
|
+
sample, ds_labels = self._get_sample_from_ar_params(ar_params)
|
|
152
|
+
|
|
153
|
+
mapper = self._create_annotation_mapper_for_detection_ar(ar_params, ds_labels)
|
|
154
|
+
|
|
155
|
+
context = self._create_detection_function_context(ar_params, sample.frame_name)
|
|
156
|
+
|
|
157
|
+
annotations = self._executor.result(
|
|
158
|
+
self._executor.submit(_worker_job_detect, context, sample.media.load_image())
|
|
159
|
+
)
|
|
160
|
+
|
|
161
|
+
tags, shapes = mapper.validate_and_remap(annotations, sample.frame_index)
|
|
162
|
+
return {"annotations": models.PatchedLabeledDataRequest(tags=tags, shapes=shapes)}
|
|
163
|
+
|
|
164
|
+
_CALCULATE_RESULT_PER_AR_TYPE = {
|
|
165
|
+
"annotate_task": _calculate_result_for_annotate_task_ar,
|
|
166
|
+
"annotate_frame": _calculate_result_for_annotate_frame_ar,
|
|
167
|
+
}
|
|
@@ -0,0 +1,214 @@
|
|
|
1
|
+
# Copyright (C) CVAT.ai Corporation
|
|
2
|
+
#
|
|
3
|
+
# SPDX-License-Identifier: MIT
|
|
4
|
+
|
|
5
|
+
from collections import OrderedDict
|
|
6
|
+
from datetime import datetime, timedelta, timezone
|
|
7
|
+
from typing import Any, Callable, TypeAlias, cast
|
|
8
|
+
|
|
9
|
+
import attrs
|
|
10
|
+
import cvat_sdk.auto_annotation as cvataa
|
|
11
|
+
import PIL.Image
|
|
12
|
+
|
|
13
|
+
from .agent_driver import (
|
|
14
|
+
AgentFunctionDriver,
|
|
15
|
+
BadArError,
|
|
16
|
+
IncompatibleFunctionError,
|
|
17
|
+
worker_current_function,
|
|
18
|
+
)
|
|
19
|
+
|
|
20
|
+
_MAX_AGE_OF_TRACKING_STATE = timedelta(hours=8)
|
|
21
|
+
|
|
22
|
+
|
|
23
|
+
@attrs.define
|
|
24
|
+
class _ExtendedTrackingState:
|
|
25
|
+
inner_state: Any # the state produced by the AA function
|
|
26
|
+
original_shape_type: str
|
|
27
|
+
original_task_id: int
|
|
28
|
+
original_image_dims: tuple[int, int]
|
|
29
|
+
last_accessed_at: datetime = attrs.field(factory=lambda: datetime.now(tz=timezone.utc))
|
|
30
|
+
|
|
31
|
+
|
|
32
|
+
class _TrackingStateContainer:
|
|
33
|
+
def __init__(self):
|
|
34
|
+
self._id_to_ext_state: OrderedDict[str, _ExtendedTrackingState] = OrderedDict()
|
|
35
|
+
|
|
36
|
+
def store(self, state: Any, shape_type: str, task_id: int, image_dims: tuple[int, int]) -> str:
|
|
37
|
+
state_id = _tracking_state_id_generator()
|
|
38
|
+
self._id_to_ext_state[state_id] = _ExtendedTrackingState(
|
|
39
|
+
inner_state=state,
|
|
40
|
+
original_shape_type=shape_type,
|
|
41
|
+
original_task_id=task_id,
|
|
42
|
+
original_image_dims=image_dims,
|
|
43
|
+
)
|
|
44
|
+
return state_id
|
|
45
|
+
|
|
46
|
+
def retrieve(self, state_id: str, task_id: int, image_dims: tuple[int, int]) -> Any:
|
|
47
|
+
ext_state = self._id_to_ext_state.get(state_id)
|
|
48
|
+
|
|
49
|
+
if not ext_state:
|
|
50
|
+
raise BadArError(f"Tracking state {state_id!r} not found - possibly expired")
|
|
51
|
+
|
|
52
|
+
if ext_state.original_task_id != task_id:
|
|
53
|
+
# This is a defense-in-depth measure. State IDs are supposed to be unguessable,
|
|
54
|
+
# but even if an attacker manages to obtain one, they will not be able to use it
|
|
55
|
+
# to get any information about a task they don't have access to.
|
|
56
|
+
raise BadArError(f"Tracking state {state_id!r} is not for task #{task_id}")
|
|
57
|
+
|
|
58
|
+
if image_dims != ext_state.original_image_dims:
|
|
59
|
+
raise BadArError("Image sizes of the start frame and the current frame are different")
|
|
60
|
+
|
|
61
|
+
ext_state.last_accessed_at = datetime.now(tz=timezone.utc)
|
|
62
|
+
self._id_to_ext_state.move_to_end(state_id)
|
|
63
|
+
|
|
64
|
+
return ext_state.inner_state, ext_state.original_shape_type
|
|
65
|
+
|
|
66
|
+
def prune(self) -> None:
|
|
67
|
+
cutoff = datetime.now(tz=timezone.utc) - _MAX_AGE_OF_TRACKING_STATE
|
|
68
|
+
|
|
69
|
+
while (
|
|
70
|
+
self._id_to_ext_state
|
|
71
|
+
and next(iter(self._id_to_ext_state.values())).last_accessed_at < cutoff
|
|
72
|
+
):
|
|
73
|
+
self._id_to_ext_state.popitem(last=False)
|
|
74
|
+
|
|
75
|
+
|
|
76
|
+
TrackingStateIdGenerator: TypeAlias = Callable[[], str]
|
|
77
|
+
|
|
78
|
+
|
|
79
|
+
_tracking_states: _TrackingStateContainer
|
|
80
|
+
_tracking_state_id_generator: TrackingStateIdGenerator
|
|
81
|
+
|
|
82
|
+
|
|
83
|
+
def _worker_job_init_tracking(
|
|
84
|
+
task_id: int,
|
|
85
|
+
image: PIL.Image.Image,
|
|
86
|
+
shapes: list[cvataa.TrackableShape],
|
|
87
|
+
) -> list[str]:
|
|
88
|
+
_tracking_states.prune()
|
|
89
|
+
|
|
90
|
+
current_function = cast(cvataa.TrackingFunction, worker_current_function())
|
|
91
|
+
|
|
92
|
+
if hasattr(current_function, "preprocess_image"):
|
|
93
|
+
pp_image = current_function.preprocess_image(_TrackingFunctionContextImpl(), image)
|
|
94
|
+
else:
|
|
95
|
+
pp_image = image
|
|
96
|
+
|
|
97
|
+
return [
|
|
98
|
+
_tracking_states.store(
|
|
99
|
+
state=current_function.init_tracking_state(
|
|
100
|
+
_TrackingFunctionShapeContextImpl(original_shape_type=shape.type), pp_image, shape
|
|
101
|
+
),
|
|
102
|
+
shape_type=shape.type,
|
|
103
|
+
task_id=task_id,
|
|
104
|
+
image_dims=image.size,
|
|
105
|
+
)
|
|
106
|
+
for shape in shapes
|
|
107
|
+
]
|
|
108
|
+
|
|
109
|
+
|
|
110
|
+
def _worker_job_track(
|
|
111
|
+
task_id: int, image: PIL.Image.Image, states: list[str]
|
|
112
|
+
) -> list[cvataa.TrackableShape | None]:
|
|
113
|
+
_tracking_states.prune()
|
|
114
|
+
|
|
115
|
+
current_function = cast(cvataa.TrackingFunction, worker_current_function())
|
|
116
|
+
|
|
117
|
+
pp_image = current_function.preprocess_image(_TrackingFunctionContextImpl(), image)
|
|
118
|
+
|
|
119
|
+
def track(state_id):
|
|
120
|
+
inner_state, original_shape_type = _tracking_states.retrieve(
|
|
121
|
+
state_id=state_id, task_id=task_id, image_dims=image.size
|
|
122
|
+
)
|
|
123
|
+
|
|
124
|
+
output_shape = current_function.track(
|
|
125
|
+
_TrackingFunctionShapeContextImpl(original_shape_type=original_shape_type),
|
|
126
|
+
pp_image,
|
|
127
|
+
inner_state,
|
|
128
|
+
)
|
|
129
|
+
|
|
130
|
+
if output_shape and output_shape.type != original_shape_type:
|
|
131
|
+
raise cvataa.BadFunctionError(
|
|
132
|
+
f"function output shape of type {output_shape.type!r}, "
|
|
133
|
+
f"but original shape was of type {original_shape_type!r}"
|
|
134
|
+
)
|
|
135
|
+
return output_shape
|
|
136
|
+
|
|
137
|
+
return list(map(track, states))
|
|
138
|
+
|
|
139
|
+
|
|
140
|
+
class _TrackingFunctionContextImpl(cvataa.TrackingFunctionContext):
|
|
141
|
+
pass
|
|
142
|
+
|
|
143
|
+
|
|
144
|
+
@attrs.frozen(kw_only=True)
|
|
145
|
+
class _TrackingFunctionShapeContextImpl(cvataa.TrackingFunctionShapeContext):
|
|
146
|
+
original_shape_type: str
|
|
147
|
+
|
|
148
|
+
|
|
149
|
+
class AgentTrackingFunctionDriver(AgentFunctionDriver):
|
|
150
|
+
FUNCTION_KIND = "tracker"
|
|
151
|
+
_function_spec: cvataa.TrackingFunctionSpec
|
|
152
|
+
|
|
153
|
+
@classmethod
|
|
154
|
+
def init_worker(cls, state_id_generator: TrackingStateIdGenerator) -> None:
|
|
155
|
+
global _tracking_states
|
|
156
|
+
_tracking_states = _TrackingStateContainer()
|
|
157
|
+
|
|
158
|
+
global _tracking_state_id_generator
|
|
159
|
+
_tracking_state_id_generator = state_id_generator
|
|
160
|
+
|
|
161
|
+
def validate_function_compatibility(self, remote_function: dict) -> None:
|
|
162
|
+
remote_supported_shape_types = frozenset(remote_function["supported_shape_types"])
|
|
163
|
+
unsupported = remote_supported_shape_types - self._function_spec.supported_shape_types
|
|
164
|
+
|
|
165
|
+
if unsupported:
|
|
166
|
+
raise IncompatibleFunctionError(
|
|
167
|
+
"the function object does not support the following shape types: "
|
|
168
|
+
+ ", ".join(map(repr, unsupported))
|
|
169
|
+
)
|
|
170
|
+
|
|
171
|
+
def _calculate_result_for_init_tracking_ar(
|
|
172
|
+
self, ar_params: dict[str, Any], check_in: object
|
|
173
|
+
) -> dict[str, Any]:
|
|
174
|
+
sample, _ = self._get_sample_from_ar_params(ar_params)
|
|
175
|
+
|
|
176
|
+
def convert_shape(shape: dict) -> cvataa.TrackableShape:
|
|
177
|
+
if shape["type"] not in self._function_spec.supported_shape_types:
|
|
178
|
+
raise BadArError(f"Unsupported shape type {shape['type']!r}")
|
|
179
|
+
return cvataa.TrackableShape(type=shape["type"], points=shape["points"])
|
|
180
|
+
|
|
181
|
+
shapes = list(map(convert_shape, ar_params["shapes"]))
|
|
182
|
+
|
|
183
|
+
states = self._executor.result(
|
|
184
|
+
self._executor.submit(
|
|
185
|
+
_worker_job_init_tracking,
|
|
186
|
+
ar_params["task"],
|
|
187
|
+
sample.media.load_image(),
|
|
188
|
+
shapes,
|
|
189
|
+
)
|
|
190
|
+
)
|
|
191
|
+
|
|
192
|
+
return {"states": states}
|
|
193
|
+
|
|
194
|
+
def _calculate_result_for_track_ar(
|
|
195
|
+
self, ar_params: dict[str, Any], check_in: object
|
|
196
|
+
) -> dict[str, Any]:
|
|
197
|
+
sample, _ = self._get_sample_from_ar_params(ar_params)
|
|
198
|
+
|
|
199
|
+
states = ar_params["states"]
|
|
200
|
+
shapes = self._executor.result(
|
|
201
|
+
self._executor.submit(
|
|
202
|
+
_worker_job_track, ar_params["task"], sample.media.load_image(), states
|
|
203
|
+
)
|
|
204
|
+
)
|
|
205
|
+
|
|
206
|
+
return {
|
|
207
|
+
"states": states,
|
|
208
|
+
"shapes": [attrs.asdict(shape) if shape else None for shape in shapes],
|
|
209
|
+
}
|
|
210
|
+
|
|
211
|
+
_CALCULATE_RESULT_PER_AR_TYPE = {
|
|
212
|
+
"init_tracking": _calculate_result_for_init_tracking_ar,
|
|
213
|
+
"track": _calculate_result_for_track_ar,
|
|
214
|
+
}
|
|
@@ -11,12 +11,9 @@ from typing import Any
|
|
|
11
11
|
import cvat_sdk.auto_annotation as cvataa
|
|
12
12
|
from cvat_sdk import Client, models
|
|
13
13
|
|
|
14
|
-
from .agent import
|
|
15
|
-
|
|
16
|
-
|
|
17
|
-
FUNCTION_PROVIDER_NATIVE,
|
|
18
|
-
run_agent,
|
|
19
|
-
)
|
|
14
|
+
from .agent import FUNCTION_PROVIDER_NATIVE, run_agent
|
|
15
|
+
from .agent_driver_detection import AgentDetectionFunctionDriver
|
|
16
|
+
from .agent_driver_tracking import AgentTrackingFunctionDriver
|
|
20
17
|
from .command_base import CommandGroup
|
|
21
18
|
from .common import FunctionLoader, configure_function_implementation_arguments
|
|
22
19
|
|
|
@@ -85,7 +82,7 @@ class FunctionCreateNative:
|
|
|
85
82
|
spec = function.spec
|
|
86
83
|
|
|
87
84
|
if isinstance(spec, cvataa.DetectionFunctionSpec):
|
|
88
|
-
remote_function["kind"] =
|
|
85
|
+
remote_function["kind"] = AgentDetectionFunctionDriver.FUNCTION_KIND
|
|
89
86
|
remote_function["labels_v2"] = []
|
|
90
87
|
|
|
91
88
|
for label_spec in spec.labels:
|
|
@@ -96,7 +93,7 @@ class FunctionCreateNative:
|
|
|
96
93
|
self._dump_sublabel_spec(sublabel) for sublabel in sublabels
|
|
97
94
|
]
|
|
98
95
|
elif isinstance(spec, cvataa.TrackingFunctionSpec):
|
|
99
|
-
remote_function["kind"] =
|
|
96
|
+
remote_function["kind"] = AgentTrackingFunctionDriver.FUNCTION_KIND
|
|
100
97
|
remote_function["supported_shape_types"] = sorted(spec.supported_shape_types)
|
|
101
98
|
else:
|
|
102
99
|
raise cvataa.BadFunctionError(f"Unsupported function spec type: {type(spec).__name__}")
|
|
@@ -0,0 +1 @@
|
|
|
1
|
+
VERSION = "2.69.0"
|
|
@@ -1,6 +1,6 @@
|
|
|
1
1
|
Metadata-Version: 2.4
|
|
2
2
|
Name: cvat-cli
|
|
3
|
-
Version: 2.
|
|
3
|
+
Version: 2.69.0
|
|
4
4
|
Summary: Command-line client for CVAT
|
|
5
5
|
Author-email: "CVAT.ai Corporation" <support@cvat.ai>
|
|
6
6
|
License-Expression: MIT
|
|
@@ -9,7 +9,7 @@ Classifier: Programming Language :: Python :: 3
|
|
|
9
9
|
Classifier: Operating System :: OS Independent
|
|
10
10
|
Requires-Python: >=3.10
|
|
11
11
|
Description-Content-Type: text/markdown
|
|
12
|
-
Requires-Dist: cvat-sdk==2.
|
|
12
|
+
Requires-Dist: cvat-sdk==2.69.0
|
|
13
13
|
Requires-Dist: attrs>=24.2.0
|
|
14
14
|
Requires-Dist: Pillow>=10.3.0
|
|
15
15
|
|
|
@@ -11,6 +11,9 @@ src/cvat_cli.egg-info/requires.txt
|
|
|
11
11
|
src/cvat_cli.egg-info/top_level.txt
|
|
12
12
|
src/cvat_cli/_internal/__init__.py
|
|
13
13
|
src/cvat_cli/_internal/agent.py
|
|
14
|
+
src/cvat_cli/_internal/agent_driver.py
|
|
15
|
+
src/cvat_cli/_internal/agent_driver_detection.py
|
|
16
|
+
src/cvat_cli/_internal/agent_driver_tracking.py
|
|
14
17
|
src/cvat_cli/_internal/command_base.py
|
|
15
18
|
src/cvat_cli/_internal/commands_all.py
|
|
16
19
|
src/cvat_cli/_internal/commands_functions.py
|
|
@@ -1 +0,0 @@
|
|
|
1
|
-
VERSION = "2.68.0"
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|