cvat-cli 2.68.0__tar.gz → 2.69.0__tar.gz

This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
Files changed (27) hide show
  1. {cvat_cli-2.68.0 → cvat_cli-2.69.0}/PKG-INFO +2 -2
  2. {cvat_cli-2.68.0 → cvat_cli-2.69.0}/pyproject.toml +1 -1
  3. {cvat_cli-2.68.0 → cvat_cli-2.69.0}/src/cvat_cli/_internal/agent.py +55 -405
  4. cvat_cli-2.69.0/src/cvat_cli/_internal/agent_driver.py +95 -0
  5. cvat_cli-2.69.0/src/cvat_cli/_internal/agent_driver_detection.py +167 -0
  6. cvat_cli-2.69.0/src/cvat_cli/_internal/agent_driver_tracking.py +214 -0
  7. {cvat_cli-2.68.0 → cvat_cli-2.69.0}/src/cvat_cli/_internal/commands_functions.py +5 -8
  8. cvat_cli-2.69.0/src/cvat_cli/version.py +1 -0
  9. {cvat_cli-2.68.0 → cvat_cli-2.69.0}/src/cvat_cli.egg-info/PKG-INFO +2 -2
  10. {cvat_cli-2.68.0 → cvat_cli-2.69.0}/src/cvat_cli.egg-info/SOURCES.txt +3 -0
  11. {cvat_cli-2.68.0 → cvat_cli-2.69.0}/src/cvat_cli.egg-info/requires.txt +1 -1
  12. cvat_cli-2.68.0/src/cvat_cli/version.py +0 -1
  13. {cvat_cli-2.68.0 → cvat_cli-2.69.0}/README.md +0 -0
  14. {cvat_cli-2.68.0 → cvat_cli-2.69.0}/setup.cfg +0 -0
  15. {cvat_cli-2.68.0 → cvat_cli-2.69.0}/src/cvat_cli/__init__.py +0 -0
  16. {cvat_cli-2.68.0 → cvat_cli-2.69.0}/src/cvat_cli/__main__.py +0 -0
  17. {cvat_cli-2.68.0 → cvat_cli-2.69.0}/src/cvat_cli/_internal/__init__.py +0 -0
  18. {cvat_cli-2.68.0 → cvat_cli-2.69.0}/src/cvat_cli/_internal/command_base.py +0 -0
  19. {cvat_cli-2.68.0 → cvat_cli-2.69.0}/src/cvat_cli/_internal/commands_all.py +0 -0
  20. {cvat_cli-2.68.0 → cvat_cli-2.69.0}/src/cvat_cli/_internal/commands_projects.py +0 -0
  21. {cvat_cli-2.68.0 → cvat_cli-2.69.0}/src/cvat_cli/_internal/commands_tasks.py +0 -0
  22. {cvat_cli-2.68.0 → cvat_cli-2.69.0}/src/cvat_cli/_internal/common.py +0 -0
  23. {cvat_cli-2.68.0 → cvat_cli-2.69.0}/src/cvat_cli/_internal/parsers.py +0 -0
  24. {cvat_cli-2.68.0 → cvat_cli-2.69.0}/src/cvat_cli/_internal/utils.py +0 -0
  25. {cvat_cli-2.68.0 → cvat_cli-2.69.0}/src/cvat_cli.egg-info/dependency_links.txt +0 -0
  26. {cvat_cli-2.68.0 → cvat_cli-2.69.0}/src/cvat_cli.egg-info/entry_points.txt +0 -0
  27. {cvat_cli-2.68.0 → cvat_cli-2.69.0}/src/cvat_cli.egg-info/top_level.txt +0 -0
@@ -1,6 +1,6 @@
1
1
  Metadata-Version: 2.4
2
2
  Name: cvat-cli
3
- Version: 2.68.0
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.68.0
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
 
@@ -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.69.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,
@@ -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
- FUNCTION_KIND_DETECTOR,
16
- FUNCTION_KIND_TRACKER,
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"] = FUNCTION_KIND_DETECTOR
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"] = FUNCTION_KIND_TRACKER
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.68.0
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.68.0
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,3 +1,3 @@
1
- cvat-sdk==2.68.0
1
+ cvat-sdk==2.69.0
2
2
  attrs>=24.2.0
3
3
  Pillow>=10.3.0
@@ -1 +0,0 @@
1
- VERSION = "2.68.0"
File without changes
File without changes