cvat-cli 2.68.0__tar.gz → 2.70.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.70.0}/PKG-INFO +2 -2
- {cvat_cli-2.68.0 → cvat_cli-2.70.0}/pyproject.toml +1 -1
- {cvat_cli-2.68.0 → cvat_cli-2.70.0}/src/cvat_cli/_internal/agent.py +55 -405
- cvat_cli-2.70.0/src/cvat_cli/_internal/agent_driver.py +116 -0
- cvat_cli-2.70.0/src/cvat_cli/_internal/agent_driver_detection.py +207 -0
- cvat_cli-2.70.0/src/cvat_cli/_internal/agent_driver_tracking.py +217 -0
- {cvat_cli-2.68.0 → cvat_cli-2.70.0}/src/cvat_cli/_internal/commands_functions.py +6 -51
- {cvat_cli-2.68.0 → cvat_cli-2.70.0}/src/cvat_cli/_internal/common.py +14 -99
- cvat_cli-2.70.0/src/cvat_cli/version.py +1 -0
- {cvat_cli-2.68.0 → cvat_cli-2.70.0}/src/cvat_cli.egg-info/PKG-INFO +2 -2
- {cvat_cli-2.68.0 → cvat_cli-2.70.0}/src/cvat_cli.egg-info/SOURCES.txt +3 -0
- {cvat_cli-2.68.0 → cvat_cli-2.70.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.70.0}/README.md +0 -0
- {cvat_cli-2.68.0 → cvat_cli-2.70.0}/setup.cfg +0 -0
- {cvat_cli-2.68.0 → cvat_cli-2.70.0}/src/cvat_cli/__init__.py +0 -0
- {cvat_cli-2.68.0 → cvat_cli-2.70.0}/src/cvat_cli/__main__.py +0 -0
- {cvat_cli-2.68.0 → cvat_cli-2.70.0}/src/cvat_cli/_internal/__init__.py +0 -0
- {cvat_cli-2.68.0 → cvat_cli-2.70.0}/src/cvat_cli/_internal/command_base.py +0 -0
- {cvat_cli-2.68.0 → cvat_cli-2.70.0}/src/cvat_cli/_internal/commands_all.py +0 -0
- {cvat_cli-2.68.0 → cvat_cli-2.70.0}/src/cvat_cli/_internal/commands_projects.py +0 -0
- {cvat_cli-2.68.0 → cvat_cli-2.70.0}/src/cvat_cli/_internal/commands_tasks.py +0 -0
- {cvat_cli-2.68.0 → cvat_cli-2.70.0}/src/cvat_cli/_internal/parsers.py +0 -0
- {cvat_cli-2.68.0 → cvat_cli-2.70.0}/src/cvat_cli/_internal/utils.py +0 -0
- {cvat_cli-2.68.0 → cvat_cli-2.70.0}/src/cvat_cli.egg-info/dependency_links.txt +0 -0
- {cvat_cli-2.68.0 → cvat_cli-2.70.0}/src/cvat_cli.egg-info/entry_points.txt +0 -0
- {cvat_cli-2.68.0 → cvat_cli-2.70.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.70.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.70.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,
|