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.
Files changed (27) hide show
  1. {cvat_cli-2.68.0 → cvat_cli-2.70.0}/PKG-INFO +2 -2
  2. {cvat_cli-2.68.0 → cvat_cli-2.70.0}/pyproject.toml +1 -1
  3. {cvat_cli-2.68.0 → cvat_cli-2.70.0}/src/cvat_cli/_internal/agent.py +55 -405
  4. cvat_cli-2.70.0/src/cvat_cli/_internal/agent_driver.py +116 -0
  5. cvat_cli-2.70.0/src/cvat_cli/_internal/agent_driver_detection.py +207 -0
  6. cvat_cli-2.70.0/src/cvat_cli/_internal/agent_driver_tracking.py +217 -0
  7. {cvat_cli-2.68.0 → cvat_cli-2.70.0}/src/cvat_cli/_internal/commands_functions.py +6 -51
  8. {cvat_cli-2.68.0 → cvat_cli-2.70.0}/src/cvat_cli/_internal/common.py +14 -99
  9. cvat_cli-2.70.0/src/cvat_cli/version.py +1 -0
  10. {cvat_cli-2.68.0 → cvat_cli-2.70.0}/src/cvat_cli.egg-info/PKG-INFO +2 -2
  11. {cvat_cli-2.68.0 → cvat_cli-2.70.0}/src/cvat_cli.egg-info/SOURCES.txt +3 -0
  12. {cvat_cli-2.68.0 → cvat_cli-2.70.0}/src/cvat_cli.egg-info/requires.txt +1 -1
  13. cvat_cli-2.68.0/src/cvat_cli/version.py +0 -1
  14. {cvat_cli-2.68.0 → cvat_cli-2.70.0}/README.md +0 -0
  15. {cvat_cli-2.68.0 → cvat_cli-2.70.0}/setup.cfg +0 -0
  16. {cvat_cli-2.68.0 → cvat_cli-2.70.0}/src/cvat_cli/__init__.py +0 -0
  17. {cvat_cli-2.68.0 → cvat_cli-2.70.0}/src/cvat_cli/__main__.py +0 -0
  18. {cvat_cli-2.68.0 → cvat_cli-2.70.0}/src/cvat_cli/_internal/__init__.py +0 -0
  19. {cvat_cli-2.68.0 → cvat_cli-2.70.0}/src/cvat_cli/_internal/command_base.py +0 -0
  20. {cvat_cli-2.68.0 → cvat_cli-2.70.0}/src/cvat_cli/_internal/commands_all.py +0 -0
  21. {cvat_cli-2.68.0 → cvat_cli-2.70.0}/src/cvat_cli/_internal/commands_projects.py +0 -0
  22. {cvat_cli-2.68.0 → cvat_cli-2.70.0}/src/cvat_cli/_internal/commands_tasks.py +0 -0
  23. {cvat_cli-2.68.0 → cvat_cli-2.70.0}/src/cvat_cli/_internal/parsers.py +0 -0
  24. {cvat_cli-2.68.0 → cvat_cli-2.70.0}/src/cvat_cli/_internal/utils.py +0 -0
  25. {cvat_cli-2.68.0 → cvat_cli-2.70.0}/src/cvat_cli.egg-info/dependency_links.txt +0 -0
  26. {cvat_cli-2.68.0 → cvat_cli-2.70.0}/src/cvat_cli.egg-info/entry_points.txt +0 -0
  27. {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.68.0
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.68.0
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
 
@@ -16,7 +16,7 @@ classifiers = [
16
16
  ]
17
17
  requires-python = ">=3.10"
18
18
  dependencies = [
19
- "cvat-sdk==2.68.0",
19
+ "cvat-sdk==2.70.0",
20
20
 
21
21
  "attrs>=24.2.0",
22
22
  "Pillow>=10.3.0",
@@ -14,35 +14,35 @@ import shutil
14
14
  import tempfile
15
15
  import threading
16
16
  import time
17
- from collections import OrderedDict
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, Any, TypeAlias
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, models
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 _RecoverableExecutor:
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
- _current_function: cvataa.AutoAnnotationFunction
120
- _tracking_states: _TrackingStateContainer
121
- _tracking_state_id_generator: _TrackingStateIdGenerator
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
- global _tracking_state_id_generator
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 _current_function.spec
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
- class _BadArError(Exception):
352
- pass
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
- class _IncompatibleFunctionError(Exception):
356
- # This should only be thrown from inside _validate_X_function_compatibility methods.
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: _RecoverableExecutor, function_id: int):
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
- self._function_spec = self._executor.result(
377
- self._executor.submit(_worker_job_get_function_spec)
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 isinstance(self._function_spec, cvataa.DetectionFunctionSpec):
426
- self._validate_detection_function_compatibility(remote_function)
427
- self._calculate_result_for_ar = self._calculate_result_for_detection_ar
428
- elif isinstance(self._function_spec, cvataa.TrackingFunctionSpec):
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
- except _IncompatibleFunctionError as ex:
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
- try:
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, _BadArError):
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
- _RecoverableExecutor(
645
+ RecoverableExecutor(
996
646
  initializer=_worker_init,
997
647
  initargs=[function_loader, _default_tracking_state_id_generator],
998
648
  ) as executor,