cvat-cli 2.67.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.67.0 → cvat_cli-2.69.0}/PKG-INFO +2 -2
- {cvat_cli-2.67.0 → cvat_cli-2.69.0}/pyproject.toml +1 -1
- {cvat_cli-2.67.0 → cvat_cli-2.69.0}/src/cvat_cli/_internal/agent.py +64 -437
- 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.67.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.67.0 → cvat_cli-2.69.0}/src/cvat_cli.egg-info/PKG-INFO +2 -2
- {cvat_cli-2.67.0 → cvat_cli-2.69.0}/src/cvat_cli.egg-info/SOURCES.txt +3 -0
- {cvat_cli-2.67.0 → cvat_cli-2.69.0}/src/cvat_cli.egg-info/requires.txt +1 -1
- cvat_cli-2.67.0/src/cvat_cli/version.py +0 -1
- {cvat_cli-2.67.0 → cvat_cli-2.69.0}/README.md +0 -0
- {cvat_cli-2.67.0 → cvat_cli-2.69.0}/setup.cfg +0 -0
- {cvat_cli-2.67.0 → cvat_cli-2.69.0}/src/cvat_cli/__init__.py +0 -0
- {cvat_cli-2.67.0 → cvat_cli-2.69.0}/src/cvat_cli/__main__.py +0 -0
- {cvat_cli-2.67.0 → cvat_cli-2.69.0}/src/cvat_cli/_internal/__init__.py +0 -0
- {cvat_cli-2.67.0 → cvat_cli-2.69.0}/src/cvat_cli/_internal/command_base.py +0 -0
- {cvat_cli-2.67.0 → cvat_cli-2.69.0}/src/cvat_cli/_internal/commands_all.py +0 -0
- {cvat_cli-2.67.0 → cvat_cli-2.69.0}/src/cvat_cli/_internal/commands_projects.py +0 -0
- {cvat_cli-2.67.0 → cvat_cli-2.69.0}/src/cvat_cli/_internal/commands_tasks.py +0 -0
- {cvat_cli-2.67.0 → cvat_cli-2.69.0}/src/cvat_cli/_internal/common.py +0 -0
- {cvat_cli-2.67.0 → cvat_cli-2.69.0}/src/cvat_cli/_internal/parsers.py +0 -0
- {cvat_cli-2.67.0 → cvat_cli-2.69.0}/src/cvat_cli/_internal/utils.py +0 -0
- {cvat_cli-2.67.0 → cvat_cli-2.69.0}/src/cvat_cli.egg-info/dependency_links.txt +0 -0
- {cvat_cli-2.67.0 → cvat_cli-2.69.0}/src/cvat_cli.egg-info/entry_points.txt +0 -0
- {cvat_cli-2.67.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(f"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
|
|
@@ -263,64 +138,41 @@ class _NewReconnectionDelay:
|
|
|
263
138
|
class _TaskCacheLimiter:
|
|
264
139
|
"""
|
|
265
140
|
This class deletes least-recently used tasks from the dataset cache,
|
|
266
|
-
so that at any time the cache contains at most
|
|
267
|
-
tasks with downloaded chunks, and at most _MAX_TASKS_WITHOUT_CHUNKS without.
|
|
141
|
+
so that at any time the cache contains at most _MAX_CACHED_TASKS tasks.
|
|
268
142
|
|
|
269
143
|
This helps manage disk usage, since agents may run indefinitely, and
|
|
270
144
|
we don't want the dataset cache to keep growing.
|
|
271
145
|
"""
|
|
272
146
|
|
|
273
|
-
|
|
274
|
-
_MAX_TASKS_WITHOUT_CHUNKS = 10
|
|
147
|
+
_MAX_CACHED_TASKS = 10
|
|
275
148
|
|
|
276
149
|
def __init__(self, client: Client) -> None:
|
|
277
150
|
self._client = client
|
|
278
151
|
self._cache_manager = make_cache_manager(client, cvatds.UpdatePolicy.IF_MISSING_OR_STALE)
|
|
279
152
|
|
|
280
|
-
self.
|
|
281
|
-
self._cached_without_chunks_task_ids = []
|
|
153
|
+
self._cached_task_ids = []
|
|
282
154
|
|
|
283
155
|
self._task_ids_in_use = set()
|
|
284
156
|
|
|
285
157
|
@contextlib.contextmanager
|
|
286
|
-
def using_cache_for_task(
|
|
287
|
-
self, task_id: int, *, with_chunks: bool
|
|
288
|
-
) -> Generator[None, None, None]:
|
|
158
|
+
def using_cache_for_task(self, task_id: int) -> Generator[None, None, None]:
|
|
289
159
|
if task_id in self._task_ids_in_use:
|
|
290
|
-
# If with_chunks is True, we would have to ensure that task_id is returned to the
|
|
291
|
-
# "with chunks" list after it leaves _task_ids_in_use, regardless of the value of
|
|
292
|
-
# with_chunks in the call that initially put it in. That would be tricky to implement,
|
|
293
|
-
# and we don't have a use case for it, so just ban it.
|
|
294
|
-
assert not with_chunks
|
|
295
160
|
yield
|
|
296
161
|
return
|
|
297
162
|
|
|
298
|
-
if task_id in self.
|
|
299
|
-
|
|
300
|
-
# _cached_with_chunks_task_ids in the end.
|
|
301
|
-
with_chunks = True
|
|
302
|
-
|
|
303
|
-
self._cached_with_chunks_task_ids.remove(task_id)
|
|
304
|
-
elif task_id in self._cached_without_chunks_task_ids:
|
|
305
|
-
self._cached_without_chunks_task_ids.remove(task_id)
|
|
163
|
+
if task_id in self._cached_task_ids:
|
|
164
|
+
self._cached_task_ids.remove(task_id)
|
|
306
165
|
|
|
307
166
|
self._task_ids_in_use.add(task_id)
|
|
308
167
|
|
|
309
|
-
if
|
|
310
|
-
|
|
311
|
-
max_cached_tasks = self._MAX_TASKS_WITH_CHUNKS
|
|
312
|
-
else:
|
|
313
|
-
cached_task_ids = self._cached_without_chunks_task_ids
|
|
314
|
-
max_cached_tasks = self._MAX_TASKS_WITHOUT_CHUNKS
|
|
315
|
-
|
|
316
|
-
if len(cached_task_ids) + len(self._task_ids_in_use) > max_cached_tasks:
|
|
317
|
-
self._delete_task_cache(cached_task_ids.pop(0))
|
|
168
|
+
if len(self._cached_task_ids) + len(self._task_ids_in_use) > self._MAX_CACHED_TASKS:
|
|
169
|
+
self._delete_task_cache(self._cached_task_ids.pop(0))
|
|
318
170
|
|
|
319
171
|
try:
|
|
320
172
|
yield
|
|
321
173
|
finally:
|
|
322
174
|
self._task_ids_in_use.remove(task_id)
|
|
323
|
-
|
|
175
|
+
self._cached_task_ids.append(task_id)
|
|
324
176
|
|
|
325
177
|
def _delete_task_cache(self, task_id: int) -> None:
|
|
326
178
|
self._client.logger.info("Deleting task %d from the cache to make room...", task_id)
|
|
@@ -371,33 +223,25 @@ def _parse_event_stream(
|
|
|
371
223
|
yield _NewReconnectionDelay(timedelta(milliseconds=int(field_value)))
|
|
372
224
|
|
|
373
225
|
|
|
374
|
-
|
|
375
|
-
|
|
376
|
-
|
|
226
|
+
def get_function_driver_class(function_spec: object) -> type[AgentFunctionDriver]:
|
|
227
|
+
if isinstance(function_spec, cvataa.DetectionFunctionSpec):
|
|
228
|
+
return AgentDetectionFunctionDriver
|
|
377
229
|
|
|
378
|
-
|
|
379
|
-
|
|
380
|
-
pass
|
|
230
|
+
if isinstance(function_spec, cvataa.TrackingFunctionSpec):
|
|
231
|
+
return AgentTrackingFunctionDriver
|
|
381
232
|
|
|
382
|
-
|
|
383
|
-
class _TrackingFunctionContextImpl(cvataa.TrackingFunctionContext):
|
|
384
|
-
pass
|
|
385
|
-
|
|
386
|
-
|
|
387
|
-
@attrs.frozen(kw_only=True)
|
|
388
|
-
class _TrackingFunctionShapeContextImpl(cvataa.TrackingFunctionShapeContext):
|
|
389
|
-
original_shape_type: str
|
|
233
|
+
raise CriticalError(f"Unsupported function spec type: {type(function_spec).__name__}")
|
|
390
234
|
|
|
391
235
|
|
|
392
236
|
class _Agent:
|
|
393
|
-
def __init__(self, client: Client, executor:
|
|
237
|
+
def __init__(self, client: Client, executor: RecoverableExecutor, function_id: int):
|
|
394
238
|
self._rng = random.Random() # nosec
|
|
395
239
|
|
|
396
240
|
self._client = client
|
|
397
|
-
self._executor = executor
|
|
398
241
|
self._function_id = function_id
|
|
399
|
-
|
|
400
|
-
|
|
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
|
|
401
245
|
)
|
|
402
246
|
|
|
403
247
|
_, response = self._client.api_client.call_api(
|
|
@@ -445,91 +289,18 @@ class _Agent:
|
|
|
445
289
|
)
|
|
446
290
|
|
|
447
291
|
try:
|
|
448
|
-
if
|
|
449
|
-
|
|
450
|
-
|
|
451
|
-
|
|
452
|
-
self._validate_tracking_function_compatibility(remote_function)
|
|
453
|
-
self._calculate_result_for_ar = self._calculate_result_for_tracking_ar
|
|
454
|
-
else:
|
|
455
|
-
raise CriticalError(
|
|
456
|
-
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})."
|
|
457
296
|
)
|
|
458
|
-
|
|
297
|
+
|
|
298
|
+
self._function_driver.validate_function_compatibility(remote_function)
|
|
299
|
+
except IncompatibleFunctionError as ex:
|
|
459
300
|
raise CriticalError(
|
|
460
301
|
f"Function #{function_id} is incompatible with function object: {ex}"
|
|
461
302
|
) from ex
|
|
462
303
|
|
|
463
|
-
def _validate_detection_function_compatibility(self, remote_function: dict) -> None:
|
|
464
|
-
self._validate_remote_function_kind(remote_function, FUNCTION_KIND_DETECTOR)
|
|
465
|
-
|
|
466
|
-
labels_by_name = {label.name: label for label in self._function_spec.labels}
|
|
467
|
-
|
|
468
|
-
for remote_label in remote_function["labels_v2"]:
|
|
469
|
-
label_desc = f"label {remote_label['name']!r}"
|
|
470
|
-
label = labels_by_name.get(remote_label["name"])
|
|
471
|
-
|
|
472
|
-
self._validate_sublabel_compatibility(remote_label, label, label_desc)
|
|
473
|
-
|
|
474
|
-
sublabels_by_name = {sl.name: sl for sl in getattr(label, "sublabels", [])}
|
|
475
|
-
|
|
476
|
-
for remote_sl in remote_label.get("sublabels", []):
|
|
477
|
-
sl_desc = f"sublabel {remote_sl['name']!r} of {label_desc}"
|
|
478
|
-
sl = sublabels_by_name.get(remote_sl["name"])
|
|
479
|
-
|
|
480
|
-
self._validate_sublabel_compatibility(remote_sl, sl, sl_desc)
|
|
481
|
-
|
|
482
|
-
def _validate_sublabel_compatibility(
|
|
483
|
-
self, remote_sl: dict, sl: models.Sublabel | None, sl_desc: str
|
|
484
|
-
):
|
|
485
|
-
if not sl:
|
|
486
|
-
raise _IncompatibleFunctionError(f"{sl_desc} is not supported.")
|
|
487
|
-
|
|
488
|
-
if remote_sl["type"] not in {"any", "unknown"} and remote_sl["type"] != sl.type:
|
|
489
|
-
raise _IncompatibleFunctionError(
|
|
490
|
-
f"{sl_desc} has type {remote_sl['type']!r}, "
|
|
491
|
-
f"but the function object declares type {sl.type!r}."
|
|
492
|
-
)
|
|
493
|
-
|
|
494
|
-
attrs_by_name = {attr.name: attr for attr in getattr(sl, "attributes", [])}
|
|
495
|
-
|
|
496
|
-
for remote_attr in remote_sl["attributes"]:
|
|
497
|
-
attr_desc = f"attribute {remote_attr['name']!r} of {sl_desc}"
|
|
498
|
-
attr = attrs_by_name.get(remote_attr["name"])
|
|
499
|
-
|
|
500
|
-
if not attr:
|
|
501
|
-
raise _IncompatibleFunctionError(f"{attr_desc} is not supported.")
|
|
502
|
-
|
|
503
|
-
if remote_attr["input_type"] != attr.input_type.value:
|
|
504
|
-
raise _IncompatibleFunctionError(
|
|
505
|
-
f"{attr_desc} has input type {remote_attr['input_type']!r},"
|
|
506
|
-
f" but the function object declares input type {attr.input_type.value!r}."
|
|
507
|
-
)
|
|
508
|
-
|
|
509
|
-
if remote_attr["values"] != attr.values:
|
|
510
|
-
raise _IncompatibleFunctionError(
|
|
511
|
-
f"{attr_desc} has values {remote_attr['values']!r},"
|
|
512
|
-
f" but the function object declares values {attr.values!r}."
|
|
513
|
-
)
|
|
514
|
-
|
|
515
|
-
def _validate_tracking_function_compatibility(self, remote_function: dict) -> None:
|
|
516
|
-
self._validate_remote_function_kind(remote_function, FUNCTION_KIND_TRACKER)
|
|
517
|
-
|
|
518
|
-
remote_supported_shape_types = frozenset(remote_function["supported_shape_types"])
|
|
519
|
-
unsupported = remote_supported_shape_types - self._function_spec.supported_shape_types
|
|
520
|
-
|
|
521
|
-
if unsupported:
|
|
522
|
-
raise _IncompatibleFunctionError(
|
|
523
|
-
"the function object does not support the following shape types: "
|
|
524
|
-
+ ", ".join(map(repr, unsupported))
|
|
525
|
-
)
|
|
526
|
-
|
|
527
|
-
def _validate_remote_function_kind(self, remote_function: dict, expected_kind: str) -> None:
|
|
528
|
-
if remote_function["kind"] != expected_kind:
|
|
529
|
-
raise _IncompatibleFunctionError(
|
|
530
|
-
f"kind is {remote_function['kind']!r} (expected {expected_kind!r})."
|
|
531
|
-
)
|
|
532
|
-
|
|
533
304
|
def _wait_between_polls(self):
|
|
534
305
|
# offset the interval randomly to avoid synchronization between workers
|
|
535
306
|
timeout_multiplier = self._rng.uniform(1 - _JITTER_AMOUNT, 1 + _JITTER_AMOUNT)
|
|
@@ -701,9 +472,23 @@ class _Agent:
|
|
|
701
472
|
)
|
|
702
473
|
self._client.logger.debug("AR %r parameters: %r", ar_id, ar_params)
|
|
703
474
|
|
|
704
|
-
|
|
705
|
-
|
|
475
|
+
last_update_timestamp = datetime.now(tz=timezone.utc)
|
|
476
|
+
|
|
477
|
+
def check_in(*, current_progress: float) -> None:
|
|
478
|
+
nonlocal last_update_timestamp
|
|
479
|
+
current_timestamp = datetime.now(tz=timezone.utc)
|
|
706
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)
|
|
707
492
|
self._complete_ar(ar_id, result)
|
|
708
493
|
except Exception as ex:
|
|
709
494
|
self._client.logger.error("Failed to process AR %r", ar_id, exc_info=True)
|
|
@@ -729,7 +514,7 @@ class _Agent:
|
|
|
729
514
|
error_message = "Failed to make an HTTP request"
|
|
730
515
|
elif isinstance(ex, cvataa.BadFunctionError):
|
|
731
516
|
error_message = "Underlying function returned incorrect result: " + str(ex)
|
|
732
|
-
elif isinstance(ex,
|
|
517
|
+
elif isinstance(ex, BadArError):
|
|
733
518
|
error_message = "Invalid annotation request: " + str(ex)
|
|
734
519
|
elif isinstance(ex, concurrent.futures.BrokenExecutor):
|
|
735
520
|
error_message = "Worker process crashed"
|
|
@@ -807,164 +592,6 @@ class _Agent:
|
|
|
807
592
|
response_data = json.loads(response.data)
|
|
808
593
|
return response_data["ar_assignment"]
|
|
809
594
|
|
|
810
|
-
def _calculate_result_for_detection_ar(self, ar_id: str, ar_params) -> dict[str, Any]:
|
|
811
|
-
if ar_params["type"] == "annotate_task":
|
|
812
|
-
with self._task_cache_limiter.using_cache_for_task(ar_params["task"], with_chunks=True):
|
|
813
|
-
return self._calculate_result_for_annotate_task_ar(ar_id, ar_params)
|
|
814
|
-
elif ar_params["type"] == "annotate_frame":
|
|
815
|
-
with self._task_cache_limiter.using_cache_for_task(
|
|
816
|
-
ar_params["task"], with_chunks=False
|
|
817
|
-
):
|
|
818
|
-
return self._calculate_result_for_annotate_frame_ar(ar_id, ar_params)
|
|
819
|
-
else:
|
|
820
|
-
raise _BadArError(f"unsupported type: {ar_params['type']!r}")
|
|
821
|
-
|
|
822
|
-
def _create_annotation_mapper_for_detection_ar(
|
|
823
|
-
self, ar_params: dict, ds_labels: Sequence[models.ILabel]
|
|
824
|
-
) -> _AnnotationMapper:
|
|
825
|
-
spec_nm = _SpecNameMapping.from_api(
|
|
826
|
-
{
|
|
827
|
-
k: models.LabelMappingEntryRequest._from_openapi_data(**v)
|
|
828
|
-
for k, v in ar_params["mapping"].items()
|
|
829
|
-
}
|
|
830
|
-
)
|
|
831
|
-
|
|
832
|
-
return _AnnotationMapper(
|
|
833
|
-
self._client.logger,
|
|
834
|
-
self._function_spec.labels,
|
|
835
|
-
ds_labels,
|
|
836
|
-
allow_unmatched_labels=False,
|
|
837
|
-
spec_nm=spec_nm,
|
|
838
|
-
conv_mask_to_poly=ar_params["conv_mask_to_poly"],
|
|
839
|
-
)
|
|
840
|
-
|
|
841
|
-
def _create_detection_function_context(
|
|
842
|
-
self, ar_params: dict, frame_name: str
|
|
843
|
-
) -> cvataa.DetectionFunctionContext:
|
|
844
|
-
return _DetectionFunctionContextImpl(
|
|
845
|
-
frame_name=frame_name,
|
|
846
|
-
conf_threshold=ar_params["threshold"],
|
|
847
|
-
conv_mask_to_poly=ar_params["conv_mask_to_poly"],
|
|
848
|
-
)
|
|
849
|
-
|
|
850
|
-
def _calculate_result_for_annotate_task_ar(self, ar_id: str, ar_params) -> dict[str, Any]:
|
|
851
|
-
ds = cvatds.TaskDataset(self._client, ar_params["task"], load_annotations=False)
|
|
852
|
-
|
|
853
|
-
# Fetching the dataset might take a while, so do a progress update to let the server
|
|
854
|
-
# know we're still alive.
|
|
855
|
-
self._update_ar(ar_id, 0)
|
|
856
|
-
last_update_timestamp = datetime.now(tz=timezone.utc)
|
|
857
|
-
|
|
858
|
-
mapper = self._create_annotation_mapper_for_detection_ar(ar_params, ds.labels)
|
|
859
|
-
|
|
860
|
-
all_annotations = models.PatchedLabeledDataRequest(tags=[], shapes=[])
|
|
861
|
-
|
|
862
|
-
for sample_index, sample in enumerate(ds.samples):
|
|
863
|
-
context = self._create_detection_function_context(ar_params, sample.frame_name)
|
|
864
|
-
annotations = self._executor.result(
|
|
865
|
-
self._executor.submit(_worker_job_detect, context, sample.media.load_image())
|
|
866
|
-
)
|
|
867
|
-
|
|
868
|
-
tags, shapes = mapper.validate_and_remap(annotations, sample.frame_index)
|
|
869
|
-
all_annotations.tags.extend(tags)
|
|
870
|
-
all_annotations.shapes.extend(shapes)
|
|
871
|
-
|
|
872
|
-
current_timestamp = datetime.now(tz=timezone.utc)
|
|
873
|
-
|
|
874
|
-
if current_timestamp >= last_update_timestamp + _UPDATE_INTERVAL:
|
|
875
|
-
self._update_ar(ar_id, (sample_index + 1) / len(ds.samples))
|
|
876
|
-
last_update_timestamp = current_timestamp
|
|
877
|
-
|
|
878
|
-
# Interactive requests are time sensitive, so if there are any,
|
|
879
|
-
# we have to put the current AR on hold and process them ASAP.
|
|
880
|
-
self._process_available_ars(REQUEST_CATEGORY_INTERACTIVE)
|
|
881
|
-
|
|
882
|
-
return {"annotations": all_annotations}
|
|
883
|
-
|
|
884
|
-
def _calculate_result_for_annotate_frame_ar(self, ar_id: str, ar_params) -> dict[str, Any]:
|
|
885
|
-
sample, ds_labels = self._get_sample_from_ar_params(ar_params)
|
|
886
|
-
|
|
887
|
-
mapper = self._create_annotation_mapper_for_detection_ar(ar_params, ds_labels)
|
|
888
|
-
|
|
889
|
-
context = self._create_detection_function_context(ar_params, sample.frame_name)
|
|
890
|
-
|
|
891
|
-
annotations = self._executor.result(
|
|
892
|
-
self._executor.submit(_worker_job_detect, context, sample.media.load_image())
|
|
893
|
-
)
|
|
894
|
-
|
|
895
|
-
tags, shapes = mapper.validate_and_remap(annotations, sample.frame_index)
|
|
896
|
-
return {"annotations": models.PatchedLabeledDataRequest(tags=tags, shapes=shapes)}
|
|
897
|
-
|
|
898
|
-
def _calculate_result_for_tracking_ar(self, ar_id: str, ar_params) -> dict[str, Any]:
|
|
899
|
-
if ar_params["type"] == "init_tracking":
|
|
900
|
-
with self._task_cache_limiter.using_cache_for_task(
|
|
901
|
-
ar_params["task"], with_chunks=False
|
|
902
|
-
):
|
|
903
|
-
return self._calculate_result_for_init_tracking_ar(ar_id, ar_params)
|
|
904
|
-
elif ar_params["type"] == "track":
|
|
905
|
-
with self._task_cache_limiter.using_cache_for_task(
|
|
906
|
-
ar_params["task"], with_chunks=False
|
|
907
|
-
):
|
|
908
|
-
return self._calculate_result_for_track_ar(ar_id, ar_params)
|
|
909
|
-
else:
|
|
910
|
-
raise _BadArError(f"unsupported type: {ar_params['type']!r}")
|
|
911
|
-
|
|
912
|
-
def _calculate_result_for_init_tracking_ar(self, ar_id: str, ar_params) -> dict[str, Any]:
|
|
913
|
-
sample, _ = self._get_sample_from_ar_params(ar_params)
|
|
914
|
-
|
|
915
|
-
def convert_shape(shape: dict) -> cvataa.TrackableShape:
|
|
916
|
-
if shape["type"] not in self._function_spec.supported_shape_types:
|
|
917
|
-
raise _BadArError(f"Unsupported shape type {shape['type']!r}")
|
|
918
|
-
return cvataa.TrackableShape(type=shape["type"], points=shape["points"])
|
|
919
|
-
|
|
920
|
-
shapes = list(map(convert_shape, ar_params["shapes"]))
|
|
921
|
-
|
|
922
|
-
states = self._executor.result(
|
|
923
|
-
self._executor.submit(
|
|
924
|
-
_worker_job_init_tracking,
|
|
925
|
-
ar_params["task"],
|
|
926
|
-
sample.media.load_image(),
|
|
927
|
-
shapes,
|
|
928
|
-
)
|
|
929
|
-
)
|
|
930
|
-
|
|
931
|
-
return {"states": states}
|
|
932
|
-
|
|
933
|
-
def _calculate_result_for_track_ar(self, ar_id: str, ar_params) -> dict[str, Any]:
|
|
934
|
-
sample, _ = self._get_sample_from_ar_params(ar_params)
|
|
935
|
-
|
|
936
|
-
states = ar_params["states"]
|
|
937
|
-
shapes = self._executor.result(
|
|
938
|
-
self._executor.submit(
|
|
939
|
-
_worker_job_track, ar_params["task"], sample.media.load_image(), states
|
|
940
|
-
)
|
|
941
|
-
)
|
|
942
|
-
|
|
943
|
-
return {
|
|
944
|
-
"states": states,
|
|
945
|
-
"shapes": [attrs.asdict(shape) if shape else None for shape in shapes],
|
|
946
|
-
}
|
|
947
|
-
|
|
948
|
-
def _get_sample_from_ar_params(self, ar_params):
|
|
949
|
-
ds = cvatds.TaskDataset(
|
|
950
|
-
self._client,
|
|
951
|
-
ar_params["task"],
|
|
952
|
-
load_annotations=False,
|
|
953
|
-
media_download_policy=cvatds.MediaDownloadPolicy.FETCH_FRAMES_ON_DEMAND,
|
|
954
|
-
)
|
|
955
|
-
|
|
956
|
-
frame_index = ar_params["frame"]
|
|
957
|
-
|
|
958
|
-
# Since ds.samples excludes deleted frames, we can't just do sample = ds.samples[frame_index].
|
|
959
|
-
# Once we drop Python 3.9, we can change this to use bisect instead of the linear search.
|
|
960
|
-
for sample in ds.samples:
|
|
961
|
-
if sample.frame_index == frame_index:
|
|
962
|
-
break
|
|
963
|
-
else:
|
|
964
|
-
raise _BadArError(f"Frame with index {frame_index} does not exist in the task")
|
|
965
|
-
|
|
966
|
-
return sample, ds.labels
|
|
967
|
-
|
|
968
595
|
def _update_ar(self, ar_id: str, progress: float) -> None:
|
|
969
596
|
self._client.logger.info("Updating AR %r progress to %.2f%%...", ar_id, progress * 100)
|
|
970
597
|
|
|
@@ -1015,7 +642,7 @@ def run_agent(
|
|
|
1015
642
|
client: Client, function_loader: FunctionLoader, function_id: int, *, burst: bool
|
|
1016
643
|
) -> None:
|
|
1017
644
|
with (
|
|
1018
|
-
|
|
645
|
+
RecoverableExecutor(
|
|
1019
646
|
initializer=_worker_init,
|
|
1020
647
|
initargs=[function_loader, _default_tracking_state_id_generator],
|
|
1021
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.67.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
|