cvat-cli 2.69.0__tar.gz → 2.71.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.69.0 → cvat_cli-2.71.0}/PKG-INFO +2 -2
- {cvat_cli-2.69.0 → cvat_cli-2.71.0}/pyproject.toml +1 -1
- {cvat_cli-2.69.0 → cvat_cli-2.71.0}/src/cvat_cli/__main__.py +15 -3
- {cvat_cli-2.69.0 → cvat_cli-2.71.0}/src/cvat_cli/_internal/agent.py +4 -0
- {cvat_cli-2.69.0 → cvat_cli-2.71.0}/src/cvat_cli/_internal/agent_driver.py +24 -3
- {cvat_cli-2.69.0 → cvat_cli-2.71.0}/src/cvat_cli/_internal/agent_driver_detection.py +44 -4
- cvat_cli-2.71.0/src/cvat_cli/_internal/agent_driver_interaction.py +218 -0
- {cvat_cli-2.69.0 → cvat_cli-2.71.0}/src/cvat_cli/_internal/agent_driver_tracking.py +5 -2
- {cvat_cli-2.69.0 → cvat_cli-2.71.0}/src/cvat_cli/_internal/command_base.py +4 -1
- {cvat_cli-2.69.0 → cvat_cli-2.71.0}/src/cvat_cli/_internal/commands_all.py +2 -0
- cvat_cli-2.71.0/src/cvat_cli/_internal/commands_config.py +45 -0
- {cvat_cli-2.69.0 → cvat_cli-2.71.0}/src/cvat_cli/_internal/commands_functions.py +5 -47
- {cvat_cli-2.69.0 → cvat_cli-2.71.0}/src/cvat_cli/_internal/common.py +15 -99
- cvat_cli-2.71.0/src/cvat_cli/version.py +1 -0
- {cvat_cli-2.69.0 → cvat_cli-2.71.0}/src/cvat_cli.egg-info/PKG-INFO +2 -2
- {cvat_cli-2.69.0 → cvat_cli-2.71.0}/src/cvat_cli.egg-info/SOURCES.txt +2 -0
- {cvat_cli-2.69.0 → cvat_cli-2.71.0}/src/cvat_cli.egg-info/requires.txt +1 -1
- cvat_cli-2.69.0/src/cvat_cli/version.py +0 -1
- {cvat_cli-2.69.0 → cvat_cli-2.71.0}/README.md +0 -0
- {cvat_cli-2.69.0 → cvat_cli-2.71.0}/setup.cfg +0 -0
- {cvat_cli-2.69.0 → cvat_cli-2.71.0}/src/cvat_cli/__init__.py +0 -0
- {cvat_cli-2.69.0 → cvat_cli-2.71.0}/src/cvat_cli/_internal/__init__.py +0 -0
- {cvat_cli-2.69.0 → cvat_cli-2.71.0}/src/cvat_cli/_internal/commands_projects.py +0 -0
- {cvat_cli-2.69.0 → cvat_cli-2.71.0}/src/cvat_cli/_internal/commands_tasks.py +0 -0
- {cvat_cli-2.69.0 → cvat_cli-2.71.0}/src/cvat_cli/_internal/parsers.py +0 -0
- {cvat_cli-2.69.0 → cvat_cli-2.71.0}/src/cvat_cli/_internal/utils.py +0 -0
- {cvat_cli-2.69.0 → cvat_cli-2.71.0}/src/cvat_cli.egg-info/dependency_links.txt +0 -0
- {cvat_cli-2.69.0 → cvat_cli-2.71.0}/src/cvat_cli.egg-info/entry_points.txt +0 -0
- {cvat_cli-2.69.0 → cvat_cli-2.71.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.71.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.71.0
|
|
13
13
|
Requires-Dist: attrs>=24.2.0
|
|
14
14
|
Requires-Dist: Pillow>=10.3.0
|
|
15
15
|
|
|
@@ -9,6 +9,7 @@ import sys
|
|
|
9
9
|
|
|
10
10
|
import urllib3.exceptions
|
|
11
11
|
from cvat_sdk import exceptions
|
|
12
|
+
from cvat_sdk.core.exceptions import AuthStoreError
|
|
12
13
|
|
|
13
14
|
from ._internal.commands_all import COMMANDS
|
|
14
15
|
from ._internal.common import (
|
|
@@ -31,10 +32,21 @@ def main(args: list[str] = None):
|
|
|
31
32
|
|
|
32
33
|
configure_logger(logger, parsed_args)
|
|
33
34
|
|
|
35
|
+
executor = popattr(parsed_args, "_executor")
|
|
36
|
+
needs_client = popattr(parsed_args, "_needs_client")
|
|
37
|
+
|
|
34
38
|
try:
|
|
35
|
-
|
|
36
|
-
|
|
37
|
-
|
|
39
|
+
if needs_client:
|
|
40
|
+
with build_client(parsed_args, logger=logger) as client:
|
|
41
|
+
executor(client, **vars(parsed_args))
|
|
42
|
+
else:
|
|
43
|
+
executor(parsed_args)
|
|
44
|
+
except (
|
|
45
|
+
exceptions.ApiException,
|
|
46
|
+
urllib3.exceptions.HTTPError,
|
|
47
|
+
CriticalError,
|
|
48
|
+
AuthStoreError,
|
|
49
|
+
) as e:
|
|
38
50
|
logger.critical(e)
|
|
39
51
|
return 1
|
|
40
52
|
|
|
@@ -36,6 +36,7 @@ from .agent_driver import (
|
|
|
36
36
|
worker_current_function,
|
|
37
37
|
)
|
|
38
38
|
from .agent_driver_detection import AgentDetectionFunctionDriver
|
|
39
|
+
from .agent_driver_interaction import AgentInteractionFunctionDriver
|
|
39
40
|
from .agent_driver_tracking import AgentTrackingFunctionDriver, TrackingStateIdGenerator
|
|
40
41
|
from .common import CriticalError, FunctionLoader
|
|
41
42
|
|
|
@@ -227,6 +228,9 @@ def get_function_driver_class(function_spec: object) -> type[AgentFunctionDriver
|
|
|
227
228
|
if isinstance(function_spec, cvataa.DetectionFunctionSpec):
|
|
228
229
|
return AgentDetectionFunctionDriver
|
|
229
230
|
|
|
231
|
+
if isinstance(function_spec, cvataa.InteractionFunctionSpec):
|
|
232
|
+
return AgentInteractionFunctionDriver
|
|
233
|
+
|
|
230
234
|
if isinstance(function_spec, cvataa.TrackingFunctionSpec):
|
|
231
235
|
return AgentTrackingFunctionDriver
|
|
232
236
|
|
|
@@ -5,7 +5,7 @@
|
|
|
5
5
|
from __future__ import annotations
|
|
6
6
|
|
|
7
7
|
from collections.abc import Callable
|
|
8
|
-
from typing import TYPE_CHECKING, Any, ClassVar, Protocol
|
|
8
|
+
from typing import TYPE_CHECKING, Any, ClassVar, Generic, Protocol, TypeVar
|
|
9
9
|
|
|
10
10
|
import cvat_sdk.auto_annotation as cvataa
|
|
11
11
|
import cvat_sdk.datasets as cvatds
|
|
@@ -47,10 +47,13 @@ class CheckInCallback(Protocol):
|
|
|
47
47
|
"""
|
|
48
48
|
|
|
49
49
|
|
|
50
|
-
|
|
50
|
+
SpecT = TypeVar("SpecT")
|
|
51
|
+
|
|
52
|
+
|
|
53
|
+
class AgentFunctionDriver(Generic[SpecT]):
|
|
51
54
|
FUNCTION_KIND: ClassVar[str]
|
|
52
55
|
|
|
53
|
-
def __init__(self, client: Client, executor: RecoverableExecutor, function_spec:
|
|
56
|
+
def __init__(self, client: Client, executor: RecoverableExecutor, function_spec: SpecT) -> None:
|
|
54
57
|
self._client = client
|
|
55
58
|
self._executor = executor
|
|
56
59
|
self._function_spec = function_spec
|
|
@@ -59,6 +62,10 @@ class AgentFunctionDriver:
|
|
|
59
62
|
def init_worker(cls, state_id_generator: TrackingStateIdGenerator) -> None:
|
|
60
63
|
pass
|
|
61
64
|
|
|
65
|
+
@classmethod
|
|
66
|
+
def get_remote_function_fields(cls, spec: SpecT) -> dict[str, Any]:
|
|
67
|
+
raise NotImplementedError
|
|
68
|
+
|
|
62
69
|
def validate_function_compatibility(self, remote_function: dict) -> None:
|
|
63
70
|
raise NotImplementedError
|
|
64
71
|
|
|
@@ -93,3 +100,17 @@ class AgentFunctionDriver:
|
|
|
93
100
|
raise BadArError(f"Frame with index {frame_index} does not exist in the task")
|
|
94
101
|
|
|
95
102
|
return sample, ds.labels
|
|
103
|
+
|
|
104
|
+
def _load_image_for_ar(self, sample, ar_params):
|
|
105
|
+
image = sample.media.load_image()
|
|
106
|
+
|
|
107
|
+
if roi := ar_params.get("roi"):
|
|
108
|
+
xtl, ytl, xbr, ybr = roi
|
|
109
|
+
width, height = image.size
|
|
110
|
+
|
|
111
|
+
if xbr > width or ybr > height:
|
|
112
|
+
raise BadArError("Invalid ROI")
|
|
113
|
+
|
|
114
|
+
return image.crop((xtl, ytl, xbr, ybr))
|
|
115
|
+
|
|
116
|
+
return image
|
|
@@ -30,9 +30,45 @@ def _worker_job_detect(
|
|
|
30
30
|
return current_function.detect(context, image)
|
|
31
31
|
|
|
32
32
|
|
|
33
|
-
class AgentDetectionFunctionDriver(AgentFunctionDriver):
|
|
33
|
+
class AgentDetectionFunctionDriver(AgentFunctionDriver[cvataa.DetectionFunctionSpec]):
|
|
34
34
|
FUNCTION_KIND = "detector"
|
|
35
|
-
|
|
35
|
+
|
|
36
|
+
@staticmethod
|
|
37
|
+
def _dump_sublabel_spec(
|
|
38
|
+
sl_spec: models.SublabelRequest | models.PatchedLabelRequest,
|
|
39
|
+
) -> dict:
|
|
40
|
+
result = {
|
|
41
|
+
"name": sl_spec.name,
|
|
42
|
+
"attributes": [
|
|
43
|
+
{
|
|
44
|
+
"name": attribute_spec.name,
|
|
45
|
+
"input_type": attribute_spec.input_type,
|
|
46
|
+
"values": attribute_spec.values,
|
|
47
|
+
}
|
|
48
|
+
for attribute_spec in getattr(sl_spec, "attributes", [])
|
|
49
|
+
],
|
|
50
|
+
}
|
|
51
|
+
|
|
52
|
+
if getattr(sl_spec, "type", "any") != "any":
|
|
53
|
+
# Add the type conditionally, to stay compatible with older
|
|
54
|
+
# CVAT versions when the function doesn't define label types.
|
|
55
|
+
result["type"] = sl_spec.type
|
|
56
|
+
|
|
57
|
+
return result
|
|
58
|
+
|
|
59
|
+
@classmethod
|
|
60
|
+
def get_remote_function_fields(cls, spec: cvataa.DetectionFunctionSpec) -> dict[str, Any]:
|
|
61
|
+
labels_v2 = []
|
|
62
|
+
|
|
63
|
+
for label_spec in spec.labels:
|
|
64
|
+
labels_v2.append(cls._dump_sublabel_spec(label_spec))
|
|
65
|
+
|
|
66
|
+
if sublabels := getattr(label_spec, "sublabels", None):
|
|
67
|
+
labels_v2[-1]["sublabels"] = [
|
|
68
|
+
cls._dump_sublabel_spec(sublabel) for sublabel in sublabels
|
|
69
|
+
]
|
|
70
|
+
|
|
71
|
+
return {"labels_v2": labels_v2}
|
|
36
72
|
|
|
37
73
|
def _validate_sublabel_compatibility(
|
|
38
74
|
self, remote_sl: dict, sl: models.Sublabel | None, sl_desc: str
|
|
@@ -134,7 +170,9 @@ class AgentDetectionFunctionDriver(AgentFunctionDriver):
|
|
|
134
170
|
for sample_index, sample in enumerate(samples):
|
|
135
171
|
context = self._create_detection_function_context(ar_params, sample.frame_name)
|
|
136
172
|
annotations = self._executor.result(
|
|
137
|
-
self._executor.submit(
|
|
173
|
+
self._executor.submit(
|
|
174
|
+
_worker_job_detect, context, self._load_image_for_ar(sample, ar_params)
|
|
175
|
+
)
|
|
138
176
|
)
|
|
139
177
|
|
|
140
178
|
tags, shapes = mapper.validate_and_remap(annotations, sample.frame_index)
|
|
@@ -155,7 +193,9 @@ class AgentDetectionFunctionDriver(AgentFunctionDriver):
|
|
|
155
193
|
context = self._create_detection_function_context(ar_params, sample.frame_name)
|
|
156
194
|
|
|
157
195
|
annotations = self._executor.result(
|
|
158
|
-
self._executor.submit(
|
|
196
|
+
self._executor.submit(
|
|
197
|
+
_worker_job_detect, context, self._load_image_for_ar(sample, ar_params)
|
|
198
|
+
)
|
|
159
199
|
)
|
|
160
200
|
|
|
161
201
|
tags, shapes = mapper.validate_and_remap(annotations, sample.frame_index)
|
|
@@ -0,0 +1,218 @@
|
|
|
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, 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_PP_IMAGE = timedelta(minutes=10)
|
|
21
|
+
|
|
22
|
+
|
|
23
|
+
class _InteractionFunctionContextImpl(cvataa.InteractionFunctionContext):
|
|
24
|
+
pass
|
|
25
|
+
|
|
26
|
+
|
|
27
|
+
@attrs.frozen(kw_only=True)
|
|
28
|
+
class _InteractionPromptsImpl(cvataa.InteractionPrompts):
|
|
29
|
+
pos_points: tuple[tuple[float, float], ...]
|
|
30
|
+
neg_points: tuple[tuple[float, float], ...]
|
|
31
|
+
bounding_box: tuple[tuple[float, float], tuple[float, float]] | None
|
|
32
|
+
|
|
33
|
+
|
|
34
|
+
@attrs.define(kw_only=True)
|
|
35
|
+
class _PpImageCacheEntry:
|
|
36
|
+
pp_image: object
|
|
37
|
+
last_accessed_at: datetime = attrs.field(factory=lambda: datetime.now(tz=timezone.utc))
|
|
38
|
+
|
|
39
|
+
|
|
40
|
+
class _PpImageCache:
|
|
41
|
+
_MISSING = object()
|
|
42
|
+
|
|
43
|
+
def __init__(self):
|
|
44
|
+
self._contents = OrderedDict[object, _PpImageCacheEntry]()
|
|
45
|
+
|
|
46
|
+
def get(self, key: object) -> object:
|
|
47
|
+
if cache_entry := self._contents.get(key):
|
|
48
|
+
cache_entry.last_accessed_at = datetime.now(tz=timezone.utc)
|
|
49
|
+
self._contents.move_to_end(key)
|
|
50
|
+
return cache_entry.pp_image
|
|
51
|
+
|
|
52
|
+
return self._MISSING
|
|
53
|
+
|
|
54
|
+
def set(self, key: object, pp_image: object) -> None:
|
|
55
|
+
if key in self._contents:
|
|
56
|
+
self._contents.move_to_end(key)
|
|
57
|
+
else:
|
|
58
|
+
self.prune()
|
|
59
|
+
|
|
60
|
+
self._contents[key] = _PpImageCacheEntry(pp_image=pp_image)
|
|
61
|
+
|
|
62
|
+
def prune(self) -> None:
|
|
63
|
+
cutoff = datetime.now(tz=timezone.utc) - _MAX_AGE_OF_PP_IMAGE
|
|
64
|
+
|
|
65
|
+
while self._contents and next(iter(self._contents.values())).last_accessed_at < cutoff:
|
|
66
|
+
self._contents.popitem(last=False)
|
|
67
|
+
|
|
68
|
+
|
|
69
|
+
_pp_image_cache: _PpImageCache
|
|
70
|
+
|
|
71
|
+
|
|
72
|
+
def _worker_job_interact(
|
|
73
|
+
cache_key: object, image: PIL.Image.Image, prompts: cvataa.InteractionPrompts
|
|
74
|
+
) -> list[cvataa.InteractionResultShape]:
|
|
75
|
+
current_function = cast(cvataa.InteractionFunction, worker_current_function())
|
|
76
|
+
context = _InteractionFunctionContextImpl()
|
|
77
|
+
|
|
78
|
+
if hasattr(current_function, "preprocess_image"):
|
|
79
|
+
pp_image = current_function.preprocess_image(context, image)
|
|
80
|
+
else:
|
|
81
|
+
pp_image = image
|
|
82
|
+
|
|
83
|
+
_pp_image_cache.set(cache_key, pp_image)
|
|
84
|
+
|
|
85
|
+
return current_function.detect(context, pp_image, prompts)
|
|
86
|
+
|
|
87
|
+
|
|
88
|
+
def _worker_job_interact_from_cache(
|
|
89
|
+
cache_key: object, prompts: cvataa.InteractionPrompts
|
|
90
|
+
) -> list[cvataa.InteractionResultShape] | None:
|
|
91
|
+
current_function = cast(cvataa.InteractionFunction, worker_current_function())
|
|
92
|
+
|
|
93
|
+
pp_image = _pp_image_cache.get(cache_key)
|
|
94
|
+
if pp_image is _PpImageCache._MISSING:
|
|
95
|
+
return None
|
|
96
|
+
|
|
97
|
+
context = _InteractionFunctionContextImpl()
|
|
98
|
+
return current_function.detect(context, pp_image, prompts)
|
|
99
|
+
|
|
100
|
+
|
|
101
|
+
class AgentInteractionFunctionDriver(AgentFunctionDriver[cvataa.InteractionFunctionSpec]):
|
|
102
|
+
FUNCTION_KIND = "interactor"
|
|
103
|
+
|
|
104
|
+
@classmethod
|
|
105
|
+
def init_worker(cls, state_id_generator: object) -> None:
|
|
106
|
+
global _pp_image_cache
|
|
107
|
+
_pp_image_cache = _PpImageCache()
|
|
108
|
+
|
|
109
|
+
@classmethod
|
|
110
|
+
def get_remote_function_fields(cls, spec: cvataa.InteractionFunctionSpec) -> dict[str, Any]:
|
|
111
|
+
fields = {"min_pos_points": spec.min_pos_points}
|
|
112
|
+
|
|
113
|
+
if spec.min_neg_points is not None:
|
|
114
|
+
fields["min_neg_points"] = spec.min_neg_points
|
|
115
|
+
|
|
116
|
+
match spec.min_bounding_boxes:
|
|
117
|
+
case 0:
|
|
118
|
+
fields["startswith_box_optional"] = True
|
|
119
|
+
case 1:
|
|
120
|
+
fields["startswith_box"] = True
|
|
121
|
+
case None:
|
|
122
|
+
pass
|
|
123
|
+
case _:
|
|
124
|
+
assert False, f"Unexpected min_bounding_boxes value: {spec.min_bounding_boxes}"
|
|
125
|
+
|
|
126
|
+
return fields
|
|
127
|
+
|
|
128
|
+
def validate_function_compatibility(self, remote_function: dict) -> None:
|
|
129
|
+
remote_min_pos_points = remote_function["min_pos_points"]
|
|
130
|
+
if remote_min_pos_points < self._function_spec.min_pos_points:
|
|
131
|
+
raise IncompatibleFunctionError(
|
|
132
|
+
"the remote function allows prompts with "
|
|
133
|
+
f"{remote_min_pos_points} positive point(s), but the function object "
|
|
134
|
+
f"requires at least {self._function_spec.min_pos_points}"
|
|
135
|
+
)
|
|
136
|
+
|
|
137
|
+
remote_min_neg_points = remote_function["min_neg_points"]
|
|
138
|
+
if remote_min_neg_points >= 0:
|
|
139
|
+
if self._function_spec.min_neg_points is None:
|
|
140
|
+
raise IncompatibleFunctionError(
|
|
141
|
+
"the remote function allows prompts with negative points, but the "
|
|
142
|
+
"function object does not"
|
|
143
|
+
)
|
|
144
|
+
|
|
145
|
+
if remote_min_neg_points < self._function_spec.min_neg_points:
|
|
146
|
+
raise IncompatibleFunctionError(
|
|
147
|
+
"the remote function allows prompts with "
|
|
148
|
+
f"{remote_min_neg_points} negative point(s), but the function object "
|
|
149
|
+
f"requires at least {self._function_spec.min_neg_points}"
|
|
150
|
+
)
|
|
151
|
+
else:
|
|
152
|
+
if self._function_spec.min_neg_points not in {None, 0}:
|
|
153
|
+
raise IncompatibleFunctionError(
|
|
154
|
+
"the remote function does not allow prompts with negative points, "
|
|
155
|
+
"but the function object requires at least "
|
|
156
|
+
f"{self._function_spec.min_neg_points}"
|
|
157
|
+
)
|
|
158
|
+
|
|
159
|
+
if self._function_spec.min_bounding_boxes is None:
|
|
160
|
+
if remote_function["startswith_box"] or remote_function["startswith_box_optional"]:
|
|
161
|
+
raise IncompatibleFunctionError(
|
|
162
|
+
"the remote function allows prompts with bounding boxes, but the "
|
|
163
|
+
"function object does not"
|
|
164
|
+
)
|
|
165
|
+
elif self._function_spec.min_bounding_boxes == 1:
|
|
166
|
+
if not remote_function["startswith_box"]:
|
|
167
|
+
raise IncompatibleFunctionError(
|
|
168
|
+
"the remote function does not require bounding boxes in prompts, "
|
|
169
|
+
"but the function object does"
|
|
170
|
+
)
|
|
171
|
+
|
|
172
|
+
def _calculate_result_for_interact_ar(
|
|
173
|
+
self, ar_params: dict[str, Any], check_in: object
|
|
174
|
+
) -> dict[str, Any]:
|
|
175
|
+
pos_points = tuple(map(tuple, ar_params["pos_points"]))
|
|
176
|
+
neg_points = tuple(map(tuple, ar_params["neg_points"]))
|
|
177
|
+
bounding_box = tuple(map(tuple, ar_params["obj_bbox"])) or None
|
|
178
|
+
|
|
179
|
+
if len(pos_points) < self._function_spec.min_pos_points:
|
|
180
|
+
raise BadArError("not enough positive points")
|
|
181
|
+
|
|
182
|
+
if self._function_spec.min_neg_points is None:
|
|
183
|
+
if len(neg_points) > 0:
|
|
184
|
+
raise BadArError("negative points are not supported")
|
|
185
|
+
else:
|
|
186
|
+
if len(neg_points) < self._function_spec.min_neg_points:
|
|
187
|
+
raise BadArError("not enough negative points")
|
|
188
|
+
|
|
189
|
+
if self._function_spec.min_bounding_boxes is None and bounding_box is not None:
|
|
190
|
+
raise BadArError("bounding boxes are not supported")
|
|
191
|
+
elif self._function_spec.min_bounding_boxes == 1 and bounding_box is None:
|
|
192
|
+
raise BadArError("a bounding box is required")
|
|
193
|
+
|
|
194
|
+
prompts = _InteractionPromptsImpl(
|
|
195
|
+
pos_points=pos_points, neg_points=neg_points, bounding_box=bounding_box
|
|
196
|
+
)
|
|
197
|
+
|
|
198
|
+
cache_key = (ar_params["task"], ar_params["frame"], tuple(ar_params.get("roi") or ()))
|
|
199
|
+
|
|
200
|
+
shapes = self._executor.result(
|
|
201
|
+
self._executor.submit(_worker_job_interact_from_cache, cache_key, prompts)
|
|
202
|
+
)
|
|
203
|
+
|
|
204
|
+
if shapes is None:
|
|
205
|
+
sample, _ = self._get_sample_from_ar_params(ar_params)
|
|
206
|
+
|
|
207
|
+
shapes = self._executor.result(
|
|
208
|
+
self._executor.submit(
|
|
209
|
+
_worker_job_interact,
|
|
210
|
+
cache_key,
|
|
211
|
+
self._load_image_for_ar(sample, ar_params),
|
|
212
|
+
prompts,
|
|
213
|
+
)
|
|
214
|
+
)
|
|
215
|
+
|
|
216
|
+
return {"shapes": [attrs.asdict(shape) for shape in shapes]}
|
|
217
|
+
|
|
218
|
+
_CALCULATE_RESULT_PER_AR_TYPE = {"interact": _calculate_result_for_interact_ar}
|
|
@@ -146,9 +146,8 @@ class _TrackingFunctionShapeContextImpl(cvataa.TrackingFunctionShapeContext):
|
|
|
146
146
|
original_shape_type: str
|
|
147
147
|
|
|
148
148
|
|
|
149
|
-
class AgentTrackingFunctionDriver(AgentFunctionDriver):
|
|
149
|
+
class AgentTrackingFunctionDriver(AgentFunctionDriver[cvataa.TrackingFunctionSpec]):
|
|
150
150
|
FUNCTION_KIND = "tracker"
|
|
151
|
-
_function_spec: cvataa.TrackingFunctionSpec
|
|
152
151
|
|
|
153
152
|
@classmethod
|
|
154
153
|
def init_worker(cls, state_id_generator: TrackingStateIdGenerator) -> None:
|
|
@@ -158,6 +157,10 @@ class AgentTrackingFunctionDriver(AgentFunctionDriver):
|
|
|
158
157
|
global _tracking_state_id_generator
|
|
159
158
|
_tracking_state_id_generator = state_id_generator
|
|
160
159
|
|
|
160
|
+
@classmethod
|
|
161
|
+
def get_remote_function_fields(cls, spec: cvataa.TrackingFunctionSpec) -> dict[str, Any]:
|
|
162
|
+
return {"supported_shape_types": sorted(spec.supported_shape_types)}
|
|
163
|
+
|
|
161
164
|
def validate_function_compatibility(self, remote_function: dict) -> None:
|
|
162
165
|
remote_supported_shape_types = frozenset(remote_function["supported_shape_types"])
|
|
163
166
|
unsupported = remote_supported_shape_types - self._function_spec.supported_shape_types
|
|
@@ -53,7 +53,10 @@ class CommandGroup:
|
|
|
53
53
|
|
|
54
54
|
for name, command in self._commands.items():
|
|
55
55
|
subparser = subparsers.add_parser(name, description=command.description)
|
|
56
|
-
subparser.set_defaults(
|
|
56
|
+
subparser.set_defaults(
|
|
57
|
+
_executor=command.execute,
|
|
58
|
+
_needs_client=getattr(command, "needs_client", True),
|
|
59
|
+
)
|
|
57
60
|
command.configure_parser(subparser)
|
|
58
61
|
|
|
59
62
|
def execute(self) -> None:
|
|
@@ -3,12 +3,14 @@
|
|
|
3
3
|
# SPDX-License-Identifier: MIT
|
|
4
4
|
|
|
5
5
|
from .command_base import CommandGroup, DeprecatedAlias
|
|
6
|
+
from .commands_config import COMMANDS as COMMANDS_CONFIG
|
|
6
7
|
from .commands_functions import COMMANDS as COMMANDS_FUNCTIONS
|
|
7
8
|
from .commands_projects import COMMANDS as COMMANDS_PROJECTS
|
|
8
9
|
from .commands_tasks import COMMANDS as COMMANDS_TASKS
|
|
9
10
|
|
|
10
11
|
COMMANDS = CommandGroup(description="Perform operations on CVAT resources.")
|
|
11
12
|
|
|
13
|
+
COMMANDS.add_command("config", COMMANDS_CONFIG)
|
|
12
14
|
COMMANDS.add_command("function", COMMANDS_FUNCTIONS)
|
|
13
15
|
COMMANDS.add_command("project", COMMANDS_PROJECTS)
|
|
14
16
|
COMMANDS.add_command("task", COMMANDS_TASKS)
|
|
@@ -0,0 +1,45 @@
|
|
|
1
|
+
# Copyright (C) CVAT.ai Corporation
|
|
2
|
+
#
|
|
3
|
+
# SPDX-License-Identifier: MIT
|
|
4
|
+
|
|
5
|
+
from __future__ import annotations
|
|
6
|
+
|
|
7
|
+
import argparse
|
|
8
|
+
|
|
9
|
+
from cvat_sdk.core.auth import AuthStore
|
|
10
|
+
|
|
11
|
+
from .command_base import CommandGroup
|
|
12
|
+
from .common import CriticalError
|
|
13
|
+
|
|
14
|
+
COMMANDS = CommandGroup(description="Manage local CVAT CLI configuration.")
|
|
15
|
+
|
|
16
|
+
|
|
17
|
+
@COMMANDS.command_class("default-server")
|
|
18
|
+
class ConfigDefaultServer:
|
|
19
|
+
needs_client = False
|
|
20
|
+
description = (
|
|
21
|
+
"Print, set, or unset the default server used by the non-profile "
|
|
22
|
+
"credential paths (--auth, CVAT_ACCESS_TOKEN, prompt)."
|
|
23
|
+
)
|
|
24
|
+
|
|
25
|
+
def configure_parser(self, parser: argparse.ArgumentParser) -> None:
|
|
26
|
+
parser.add_argument("server", nargs="?", default=None, help="server URL to remember")
|
|
27
|
+
parser.add_argument("--unset", action="store_true", help="clear the saved default server")
|
|
28
|
+
|
|
29
|
+
def execute(self, args: argparse.Namespace) -> None:
|
|
30
|
+
store = AuthStore()
|
|
31
|
+
|
|
32
|
+
if args.unset and args.server is not None:
|
|
33
|
+
raise CriticalError("Cannot combine a server value with --unset.")
|
|
34
|
+
|
|
35
|
+
if args.unset:
|
|
36
|
+
store.clear_default_server()
|
|
37
|
+
print("Default server cleared.")
|
|
38
|
+
elif args.server is not None:
|
|
39
|
+
if not args.server.strip():
|
|
40
|
+
raise CriticalError("Default server cannot be empty. Use --unset to clear it.")
|
|
41
|
+
store.set_default_server(args.server)
|
|
42
|
+
print(f"Default server is now {args.server!r}.")
|
|
43
|
+
else:
|
|
44
|
+
current = store.get_default_server()
|
|
45
|
+
print(current if current is not None else "(no default server set)")
|
|
@@ -8,12 +8,9 @@ import textwrap
|
|
|
8
8
|
from collections.abc import Sequence
|
|
9
9
|
from typing import Any
|
|
10
10
|
|
|
11
|
-
|
|
12
|
-
from cvat_sdk import Client, models
|
|
11
|
+
from cvat_sdk import Client
|
|
13
12
|
|
|
14
|
-
from .agent import FUNCTION_PROVIDER_NATIVE, run_agent
|
|
15
|
-
from .agent_driver_detection import AgentDetectionFunctionDriver
|
|
16
|
-
from .agent_driver_tracking import AgentTrackingFunctionDriver
|
|
13
|
+
from .agent import FUNCTION_PROVIDER_NATIVE, get_function_driver_class, run_agent
|
|
17
14
|
from .command_base import CommandGroup
|
|
18
15
|
from .common import FunctionLoader, configure_function_implementation_arguments
|
|
19
16
|
|
|
@@ -40,29 +37,6 @@ class FunctionCreateNative:
|
|
|
40
37
|
|
|
41
38
|
configure_function_implementation_arguments(parser)
|
|
42
39
|
|
|
43
|
-
@staticmethod
|
|
44
|
-
def _dump_sublabel_spec(
|
|
45
|
-
sl_spec: models.SublabelRequest | models.PatchedLabelRequest,
|
|
46
|
-
) -> dict:
|
|
47
|
-
result = {
|
|
48
|
-
"name": sl_spec.name,
|
|
49
|
-
"attributes": [
|
|
50
|
-
{
|
|
51
|
-
"name": attribute_spec.name,
|
|
52
|
-
"input_type": attribute_spec.input_type,
|
|
53
|
-
"values": attribute_spec.values,
|
|
54
|
-
}
|
|
55
|
-
for attribute_spec in getattr(sl_spec, "attributes", [])
|
|
56
|
-
],
|
|
57
|
-
}
|
|
58
|
-
|
|
59
|
-
if getattr(sl_spec, "type", "any") != "any":
|
|
60
|
-
# Add the type conditionally, to stay compatible with older
|
|
61
|
-
# CVAT versions when the function doesn't define label types.
|
|
62
|
-
result["type"] = sl_spec.type
|
|
63
|
-
|
|
64
|
-
return result
|
|
65
|
-
|
|
66
40
|
def execute(
|
|
67
41
|
self,
|
|
68
42
|
client: Client,
|
|
@@ -72,32 +46,16 @@ class FunctionCreateNative:
|
|
|
72
46
|
function_loader: FunctionLoader,
|
|
73
47
|
) -> None:
|
|
74
48
|
function = function_loader.load()
|
|
49
|
+
driver_class = get_function_driver_class(function.spec)
|
|
75
50
|
|
|
76
51
|
remote_function: dict[str, Any] = {
|
|
77
52
|
"provider": FUNCTION_PROVIDER_NATIVE,
|
|
78
53
|
"name": name,
|
|
79
54
|
"visibility": visibility,
|
|
55
|
+
"kind": driver_class.FUNCTION_KIND,
|
|
56
|
+
**driver_class.get_remote_function_fields(function.spec),
|
|
80
57
|
}
|
|
81
58
|
|
|
82
|
-
spec = function.spec
|
|
83
|
-
|
|
84
|
-
if isinstance(spec, cvataa.DetectionFunctionSpec):
|
|
85
|
-
remote_function["kind"] = AgentDetectionFunctionDriver.FUNCTION_KIND
|
|
86
|
-
remote_function["labels_v2"] = []
|
|
87
|
-
|
|
88
|
-
for label_spec in spec.labels:
|
|
89
|
-
remote_function["labels_v2"].append(self._dump_sublabel_spec(label_spec))
|
|
90
|
-
|
|
91
|
-
if sublabels := getattr(label_spec, "sublabels", None):
|
|
92
|
-
remote_function["labels_v2"][-1]["sublabels"] = [
|
|
93
|
-
self._dump_sublabel_spec(sublabel) for sublabel in sublabels
|
|
94
|
-
]
|
|
95
|
-
elif isinstance(spec, cvataa.TrackingFunctionSpec):
|
|
96
|
-
remote_function["kind"] = AgentTrackingFunctionDriver.FUNCTION_KIND
|
|
97
|
-
remote_function["supported_shape_types"] = sorted(spec.supported_shape_types)
|
|
98
|
-
else:
|
|
99
|
-
raise cvataa.BadFunctionError(f"Unsupported function spec type: {type(spec).__name__}")
|
|
100
|
-
|
|
101
59
|
_, response = client.api_client.call_api(
|
|
102
60
|
"/api/functions",
|
|
103
61
|
"POST",
|
|
@@ -4,27 +4,23 @@
|
|
|
4
4
|
# SPDX-License-Identifier: MIT
|
|
5
5
|
|
|
6
6
|
import argparse
|
|
7
|
-
import getpass
|
|
8
7
|
import importlib
|
|
9
8
|
import importlib.util
|
|
10
9
|
import logging
|
|
11
|
-
import os
|
|
12
10
|
import sys
|
|
13
|
-
import textwrap
|
|
14
|
-
from collections.abc import Callable
|
|
15
11
|
from http.client import HTTPConnection
|
|
16
12
|
from pathlib import Path
|
|
17
13
|
from typing import Any
|
|
18
14
|
|
|
19
15
|
import attrs
|
|
20
16
|
import cvat_sdk.auto_annotation as cvataa
|
|
21
|
-
from cvat_sdk.core.
|
|
22
|
-
|
|
23
|
-
|
|
24
|
-
|
|
25
|
-
Credentials,
|
|
26
|
-
PasswordCredentials,
|
|
17
|
+
from cvat_sdk.core.auth import (
|
|
18
|
+
ClientAuthParameters,
|
|
19
|
+
configure_client_auth_arguments,
|
|
20
|
+
make_client_from_cli,
|
|
27
21
|
)
|
|
22
|
+
from cvat_sdk.core.client import Client
|
|
23
|
+
from cvat_sdk.core.exceptions import AuthStoreError
|
|
28
24
|
|
|
29
25
|
from ..version import VERSION
|
|
30
26
|
from .parsers import BuildDictAction, parse_function_parameter
|
|
@@ -35,82 +31,9 @@ class CriticalError(Exception):
|
|
|
35
31
|
pass
|
|
36
32
|
|
|
37
33
|
|
|
38
|
-
CVAT_ACCESS_TOKEN_ENV_VAR = "CVAT_ACCESS_TOKEN" # nosec - a variable name declaration
|
|
39
|
-
|
|
40
|
-
|
|
41
|
-
def default_auth_factory() -> Callable[[str], Credentials]:
|
|
42
|
-
"""
|
|
43
|
-
Try to read the CVAT_ACCESS_TOKEN environment variable for a Personal Access Token (PAT).
|
|
44
|
-
If there is no value, try using the current user and asking for the password.
|
|
45
|
-
"""
|
|
46
|
-
|
|
47
|
-
token = os.getenv(CVAT_ACCESS_TOKEN_ENV_VAR)
|
|
48
|
-
if token is not None:
|
|
49
|
-
return lambda _: AccessTokenCredentials(token)
|
|
50
|
-
|
|
51
|
-
return get_auth_factory(getpass.getuser())
|
|
52
|
-
|
|
53
|
-
|
|
54
|
-
def get_auth_factory(s: str) -> Callable[[str], Credentials]:
|
|
55
|
-
"""
|
|
56
|
-
Parse a USER[:PASS] string and return a callable that takes the server URL
|
|
57
|
-
and returns auth credentials for that URL.
|
|
58
|
-
The callable will prompt the user for the password if none was initially supplied in the
|
|
59
|
-
input string and in the PASS env variable.
|
|
60
|
-
"""
|
|
61
|
-
|
|
62
|
-
user, _, password = s.partition(":")
|
|
63
|
-
if not password:
|
|
64
|
-
password = os.environ.get("PASS")
|
|
65
|
-
|
|
66
|
-
if password:
|
|
67
|
-
return lambda _: PasswordCredentials(user, password)
|
|
68
|
-
else:
|
|
69
|
-
return lambda url: PasswordCredentials(
|
|
70
|
-
user, getpass.getpass(f"Password for {user} at {url}: ")
|
|
71
|
-
)
|
|
72
|
-
|
|
73
|
-
|
|
74
34
|
def configure_common_arguments(parser: argparse.ArgumentParser) -> None:
|
|
75
35
|
parser.add_argument("--version", action="version", version=VERSION)
|
|
76
|
-
parser
|
|
77
|
-
"--insecure",
|
|
78
|
-
action="store_true",
|
|
79
|
-
help="Allows to disable SSL certificate check",
|
|
80
|
-
)
|
|
81
|
-
|
|
82
|
-
parser.add_argument(
|
|
83
|
-
"--auth",
|
|
84
|
-
type=get_auth_factory,
|
|
85
|
-
metavar="USER[:PASS]",
|
|
86
|
-
default=default_auth_factory(),
|
|
87
|
-
help=textwrap.dedent("""\
|
|
88
|
-
User and password to use for authentication;
|
|
89
|
-
defaults to the current user and supports the PASS
|
|
90
|
-
environment variable or password prompt.
|
|
91
|
-
A Personal Access Token (PAT) can be generated on the server
|
|
92
|
-
and specified in the {} environment variable instead.
|
|
93
|
-
(default user: {}).
|
|
94
|
-
""").format(CVAT_ACCESS_TOKEN_ENV_VAR, getpass.getuser()),
|
|
95
|
-
)
|
|
96
|
-
parser.add_argument(
|
|
97
|
-
"--server-host", type=str, default="http://localhost", help="host (default: %(default)s)"
|
|
98
|
-
)
|
|
99
|
-
parser.add_argument(
|
|
100
|
-
"--server-port",
|
|
101
|
-
type=int,
|
|
102
|
-
default=None,
|
|
103
|
-
help="port (default: 80 for http and 443 for https connections)",
|
|
104
|
-
)
|
|
105
|
-
parser.add_argument(
|
|
106
|
-
"--organization",
|
|
107
|
-
"--org",
|
|
108
|
-
metavar="SLUG",
|
|
109
|
-
help="""short name (slug) of the organization
|
|
110
|
-
to use when listing or creating resources;
|
|
111
|
-
set to blank string to use the personal workspace
|
|
112
|
-
(default: list all accessible objects, create in personal workspace)""",
|
|
113
|
-
)
|
|
36
|
+
configure_client_auth_arguments(parser)
|
|
114
37
|
parser.add_argument(
|
|
115
38
|
"--debug",
|
|
116
39
|
action="store_const",
|
|
@@ -119,6 +42,7 @@ def configure_common_arguments(parser: argparse.ArgumentParser) -> None:
|
|
|
119
42
|
default=logging.INFO,
|
|
120
43
|
help="show debug output",
|
|
121
44
|
)
|
|
45
|
+
parser.set_defaults(_needs_client=True)
|
|
122
46
|
|
|
123
47
|
|
|
124
48
|
def configure_logger(logger: logging.Logger, parsed_args: argparse.Namespace) -> None:
|
|
@@ -135,24 +59,16 @@ def configure_logger(logger: logging.Logger, parsed_args: argparse.Namespace) ->
|
|
|
135
59
|
|
|
136
60
|
|
|
137
61
|
def build_client(parsed_args: argparse.Namespace, logger: logging.Logger) -> Client:
|
|
138
|
-
|
|
139
|
-
|
|
140
|
-
|
|
141
|
-
if server_port := popattr(parsed_args, "server_port"):
|
|
142
|
-
url += f":{server_port}"
|
|
62
|
+
auth_args = ClientAuthParameters.from_namespace(parsed_args)
|
|
63
|
+
for field in attrs.fields(ClientAuthParameters):
|
|
64
|
+
popattr(parsed_args, field.name)
|
|
143
65
|
|
|
144
|
-
|
|
145
|
-
|
|
146
|
-
|
|
147
|
-
|
|
148
|
-
check_server_version=False, # version is checked after auth to support versions < 2.3
|
|
149
|
-
)
|
|
66
|
+
try:
|
|
67
|
+
client = make_client_from_cli(auth_args, logger=logger)
|
|
68
|
+
except AuthStoreError as e:
|
|
69
|
+
raise CriticalError(str(e)) from e
|
|
150
70
|
|
|
151
|
-
client.login(popattr(parsed_args, "auth")(client.api_client.configuration.host))
|
|
152
71
|
client.check_server_version(fail_if_unsupported=False)
|
|
153
|
-
|
|
154
|
-
client.organization_slug = popattr(parsed_args, "organization")
|
|
155
|
-
|
|
156
72
|
return client
|
|
157
73
|
|
|
158
74
|
|
|
@@ -0,0 +1 @@
|
|
|
1
|
+
VERSION = "2.71.0"
|
|
@@ -1,6 +1,6 @@
|
|
|
1
1
|
Metadata-Version: 2.4
|
|
2
2
|
Name: cvat-cli
|
|
3
|
-
Version: 2.
|
|
3
|
+
Version: 2.71.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.71.0
|
|
13
13
|
Requires-Dist: attrs>=24.2.0
|
|
14
14
|
Requires-Dist: Pillow>=10.3.0
|
|
15
15
|
|
|
@@ -13,9 +13,11 @@ src/cvat_cli/_internal/__init__.py
|
|
|
13
13
|
src/cvat_cli/_internal/agent.py
|
|
14
14
|
src/cvat_cli/_internal/agent_driver.py
|
|
15
15
|
src/cvat_cli/_internal/agent_driver_detection.py
|
|
16
|
+
src/cvat_cli/_internal/agent_driver_interaction.py
|
|
16
17
|
src/cvat_cli/_internal/agent_driver_tracking.py
|
|
17
18
|
src/cvat_cli/_internal/command_base.py
|
|
18
19
|
src/cvat_cli/_internal/commands_all.py
|
|
20
|
+
src/cvat_cli/_internal/commands_config.py
|
|
19
21
|
src/cvat_cli/_internal/commands_functions.py
|
|
20
22
|
src/cvat_cli/_internal/commands_projects.py
|
|
21
23
|
src/cvat_cli/_internal/commands_tasks.py
|
|
@@ -1 +0,0 @@
|
|
|
1
|
-
VERSION = "2.69.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
|