opensportslib 0.3.1.dev2__py3-none-any.whl → 0.3.1.dev4__py3-none-any.whl

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 (33) hide show
  1. opensportslib/apis/__init__.py +2 -0
  2. opensportslib/apis/base_task_model.py +39 -8
  3. opensportslib/apis/classification.py +35 -3
  4. opensportslib/apis/config.py +5 -0
  5. opensportslib/apis/configuration.py +145 -0
  6. opensportslib/apis/localization.py +33 -3
  7. opensportslib/apis/vqa.py +7 -3
  8. opensportslib/core/config/editable.py +396 -0
  9. opensportslib/core/config/rule_variants.py +66 -0
  10. opensportslib/core/trainer/vqa_trainer.py +10 -0
  11. opensportslib/core/utils/direct_video.py +66 -0
  12. opensportslib/metrics/classification_metric.py +9 -11
  13. opensportslib/models/base/qwen_vl_native.py +4 -2
  14. opensportslib/models/base/qwen_xvars.py +4 -2
  15. opensportslib/models/base/rule_based.py +1 -62
  16. opensportslib/models/base/xvars_videochatgpt.py +6 -3
  17. opensportslib/setup/setup.py +1 -1
  18. opensportslib/tools/hf_transfer.py +135 -20
  19. {opensportslib-0.3.1.dev2.dist-info → opensportslib-0.3.1.dev4.dist-info}/METADATA +90 -1
  20. {opensportslib-0.3.1.dev2.dist-info → opensportslib-0.3.1.dev4.dist-info}/RECORD +33 -26
  21. tests/conftest.py +2 -0
  22. tests/test_config_architecture.py +4 -4
  23. tests/test_editable_config.py +190 -0
  24. tests/test_hf_transfer_tools.py +147 -0
  25. tests/test_localization_hf_backend_override.py +1 -0
  26. tests/test_optional_hf_config.py +84 -0
  27. tests/test_setup_cli.py +4 -4
  28. tests/test_vqa_xvars_videochatgpt.py +4 -8
  29. {opensportslib-0.3.1.dev2.dist-info → opensportslib-0.3.1.dev4.dist-info}/WHEEL +0 -0
  30. {opensportslib-0.3.1.dev2.dist-info → opensportslib-0.3.1.dev4.dist-info}/entry_points.txt +0 -0
  31. {opensportslib-0.3.1.dev2.dist-info → opensportslib-0.3.1.dev4.dist-info}/licenses/LICENSE +0 -0
  32. {opensportslib-0.3.1.dev2.dist-info → opensportslib-0.3.1.dev4.dist-info}/licenses/LICENSE-COMMERCIAL +0 -0
  33. {opensportslib-0.3.1.dev2.dist-info → opensportslib-0.3.1.dev4.dist-info}/top_level.txt +0 -0
@@ -1,6 +1,7 @@
1
1
  # opensportslib/apis/__init__.py
2
2
 
3
3
  # Import task APIs
4
+ from opensportslib.apis.config import Config
4
5
  from opensportslib.apis.base_task_model import BaseTaskModel
5
6
  from opensportslib.apis.classification import ClassificationModel
6
7
  from opensportslib.apis.localization import LocalizationModel
@@ -10,6 +11,7 @@ warnings.filterwarnings("ignore")
10
11
 
11
12
  # Expose only these
12
13
  __all__ = [
14
+ "Config",
13
15
  "BaseTaskModel",
14
16
  "ClassificationModel",
15
17
  "LocalizationModel",
@@ -16,10 +16,12 @@ from typing import Any
16
16
  from urllib import error as urlerror
17
17
  from urllib import request as urlrequest
18
18
 
19
+ from opensportslib.apis.configuration import ConfigurationMixin
19
20
  from opensportslib.core.config.accessors import get_component_name_by_kind
21
+ from opensportslib.core.config.editable import Config
22
+ from opensportslib.core.config.runtime_adapter import dict_to_namespace
20
23
  from opensportslib.core.utils.config import (
21
24
  expand,
22
- load_config_omega,
23
25
  fetch_and_merge_config_from_HF,
24
26
  resolve_config_path,
25
27
  resolve_inference_class_metadata,
@@ -55,8 +57,7 @@ def _manifest_media_references(payload: dict[str, Any]):
55
57
  if isinstance(value, str):
56
58
  yield value, lambda replacement, values=values, index=index: values.__setitem__(index, replacement)
57
59
 
58
-
59
- class BaseTaskModel(ABC):
60
+ class BaseTaskModel(ConfigurationMixin, ABC):
60
61
  """Thin shared contract for task-level OpenSportsLib wrappers."""
61
62
 
62
63
  def __init__(
@@ -78,11 +79,38 @@ class BaseTaskModel(ABC):
78
79
  if self.remote_timeout <= 0 or self.remote_poll_interval <= 0 or self.remote_result_timeout <= 0:
79
80
  raise ValueError("Remote timeout values must be positive.")
80
81
 
81
- if config is None:
82
- raise ValueError("config path is required")
82
+ self._config_editor = copy.deepcopy(config) if isinstance(config, Config) else None
83
+ if self._config_editor is not None:
84
+ weights = weights if weights is not None else self._config_editor.weights
85
+ config = self._config_editor.source
86
+ if self.is_remote:
87
+ self._config_editor.remote_overrides()
88
+
89
+ if config is None and self._config_editor is None:
90
+ from huggingface_hub.utils import HFValidationError, validate_repo_id
83
91
 
84
- self.config_path = resolve_config_path(config)
85
- self.config = load_config_omega(self.config_path)
92
+ if (not isinstance(weights, str) or os.path.exists(expand(weights))
93
+ or weights.endswith((".pt", ".pth", ".tar"))):
94
+ raise ValueError("config path is required unless weights is a Hugging Face model ID")
95
+ try:
96
+ validate_repo_id(weights)
97
+ except HFValidationError as exc:
98
+ raise ValueError("config path is required unless weights is a Hugging Face model ID") from exc
99
+ try:
100
+ config = resolve_config_path(weights)
101
+ except Exception as exc:
102
+ raise ValueError(
103
+ f"Could not load OpenSportsLib config.yaml from {weights!r}; "
104
+ "provide config explicitly or publish a compatible config.yaml."
105
+ ) from exc
106
+
107
+ if self._config_editor is not None:
108
+ self.config_path = self._config_editor.source
109
+ self.config = dict_to_namespace(self._config_editor.get_config())
110
+ else:
111
+ self.config_path = resolve_config_path(config)
112
+ self._config_editor = Config.from_file(self.config_path)
113
+ self.config = dict_to_namespace(self._config_editor.get_config())
86
114
  self.last_loaded_weights = None
87
115
  self.best_checkpoint = None
88
116
 
@@ -96,6 +124,7 @@ class BaseTaskModel(ABC):
96
124
  self.last_loaded_weights = weights
97
125
  self.best_checkpoint = weights
98
126
 
127
+ self.config = self._effective_config(self.config)
99
128
  self.train_flag = False # Flag to indicate whether we're in training mode (affects checkpoint loading behavior)
100
129
 
101
130
  data_cfg = getattr(self.config, "DATA", None)
@@ -253,7 +282,7 @@ class BaseTaskModel(ABC):
253
282
  model_id: str | None = None,
254
283
  task_options: dict[str, Any] | None = None,
255
284
  ) -> dict[str, Any]:
256
- """Submit a single uploaded video, primarily for direct VQA inference."""
285
+ """Submit one video for direct task inference."""
257
286
 
258
287
  if not self.remote:
259
288
  raise RuntimeError("Remote inference is not configured. Pass `remote=` to the model constructor.")
@@ -405,6 +434,8 @@ class BaseTaskModel(ABC):
405
434
  return self._open_request(request)
406
435
 
407
436
  def _post_multipart(self, endpoint: str, *, fields: dict[str, str], files: dict[str, Path]) -> dict[str, Any]:
437
+ if endpoint == "/predict":
438
+ fields = self._config_request_fields(fields)
408
439
  boundary = f"----OpenSportsLib{uuid.uuid4().hex}"
409
440
  chunks: list[bytes] = []
410
441
  for name, value in fields.items():
@@ -7,6 +7,7 @@ import os
7
7
  import json
8
8
 
9
9
  from opensportslib.apis.base_task_model import BaseTaskModel
10
+ from opensportslib.apis.configuration import config_operation
10
11
  from opensportslib.core.config.accessors import (
11
12
  get_component_provider_by_kind,
12
13
  get_data_modality,
@@ -180,6 +181,7 @@ class ClassificationModel(BaseTaskModel):
180
181
  # public training interface
181
182
  # -----------------------------------------------------------------
182
183
 
184
+ @config_operation
183
185
  def train(
184
186
  self,
185
187
  train_set=None,
@@ -200,7 +202,7 @@ class ClassificationModel(BaseTaskModel):
200
202
  train_set = self._resolve_split_path("train", train_set)
201
203
  valid_set = self._resolve_split_path("valid", valid_set)
202
204
 
203
- self.config = resolve_config_omega(self.config, weights=weights)
205
+ self.config = self._effective_config(resolve_config_omega(self.config, weights=weights))
204
206
  logging.info("Configuration:")
205
207
  logging.info(self.config)
206
208
 
@@ -251,15 +253,19 @@ class ClassificationModel(BaseTaskModel):
251
253
  self.last_loaded_weights = self.best_checkpoint
252
254
  return self.best_checkpoint
253
255
 
256
+ @config_operation
254
257
  def infer(
255
258
  self,
256
259
  test_set=None,
257
260
  weights=None,
258
261
  use_ddp=False,
259
262
  use_wandb=True,
263
+ video_path: str | None = None,
260
264
  **kwargs,
261
265
  ):
262
266
  """Run model inference and return predictions in OSL JSON format."""
267
+ if test_set is not None and video_path is not None:
268
+ raise ValueError("Provide either `test_set` or `video_path`, not both.")
263
269
  remote_mode_provided = "remote_mode" in kwargs
264
270
  remote_mode = kwargs.pop("remote_mode", "full_test_set")
265
271
  if self.is_remote:
@@ -267,6 +273,17 @@ class ClassificationModel(BaseTaskModel):
267
273
  remote_task_options = kwargs.pop("remote_task_options", None)
268
274
  if kwargs:
269
275
  raise TypeError(f"Unsupported remote inference options: {', '.join(kwargs)}")
276
+ if video_path is not None:
277
+ if remote_mode != "full_test_set":
278
+ raise ValueError("`remote_mode=per_sample` requires `test_set`, not direct video input.")
279
+ job = self.submit_video_inference(
280
+ task_type="classification",
281
+ video_path=video_path,
282
+ model_id=remote_model_id,
283
+ task_options=remote_task_options,
284
+ )
285
+ self.last_remote_failures = []
286
+ return self.wait_for_remote_result(job["job_id"])["result"]["predictions"]
270
287
  test_set = self._resolve_split_path("test", test_set)
271
288
  if remote_mode == "per_sample":
272
289
  batch = self.submit_per_sample_inference(
@@ -294,13 +311,27 @@ class ClassificationModel(BaseTaskModel):
294
311
  raise ValueError("`remote_mode` is available only when `remote` is configured.")
295
312
  del kwargs
296
313
 
314
+ if video_path is not None:
315
+ from opensportslib.core.utils.direct_video import direct_video_manifest
316
+ from opensportslib.core.utils.config import resolve_config_omega
317
+
318
+ manifest_config = self._effective_config(resolve_config_omega(self.config, weights=weights))
319
+ manifest_config = resolve_inference_class_metadata(manifest_config)
320
+ with direct_video_manifest(manifest_config, video_path, "classification") as manifest:
321
+ return self.infer(
322
+ test_set=manifest,
323
+ weights=weights,
324
+ use_ddp=use_ddp,
325
+ use_wandb=use_wandb,
326
+ )
327
+
297
328
  import torch
298
329
  import torch.multiprocessing as mp
299
330
  from opensportslib.core.utils.config import resolve_config_omega
300
331
 
301
332
  test_set = self._resolve_split_path("test", test_set)
302
333
 
303
- self.config = resolve_config_omega(self.config, weights=weights)
334
+ self.config = self._effective_config(resolve_config_omega(self.config, weights=weights))
304
335
  self.config = resolve_inference_class_metadata(self.config)
305
336
  logging.info("Configuration:")
306
337
  logging.info(self.config)
@@ -349,6 +380,7 @@ class ClassificationModel(BaseTaskModel):
349
380
  predictions = json.load(f)
350
381
  return predictions
351
382
 
383
+ @config_operation
352
384
  def evaluate(
353
385
  self,
354
386
  test_set=None,
@@ -367,7 +399,7 @@ class ClassificationModel(BaseTaskModel):
367
399
 
368
400
  test_set = self._resolve_split_path("test", test_set)
369
401
 
370
- self.config = resolve_config_omega(self.config, weights=weights)
402
+ self.config = self._effective_config(resolve_config_omega(self.config, weights=weights))
371
403
  self.config = resolve_inference_class_metadata(self.config)
372
404
  logging.info("Configuration:")
373
405
  logging.info(self.config)
@@ -0,0 +1,5 @@
1
+ """Public configuration API."""
2
+
3
+ from opensportslib.core.config.editable import Config
4
+
5
+ __all__ = ["Config"]
@@ -0,0 +1,145 @@
1
+ """Configuration lifecycle shared by task wrappers."""
2
+
3
+ from copy import deepcopy
4
+ from functools import wraps
5
+ import inspect
6
+ import os
7
+ import json
8
+ from threading import RLock
9
+ from urllib.parse import urlencode
10
+ from urllib.request import Request
11
+
12
+ from opensportslib.core.config.editable import Config, _get, _set
13
+ from opensportslib.core.config.runtime_adapter import dict_to_namespace, namespace_to_plain_dict
14
+
15
+
16
+ def config_operation(method):
17
+ """Serialize operations and prevent edits while an operation is active."""
18
+ @wraps(method)
19
+ def wrapped(self, *args, **kwargs):
20
+ lock = self.__dict__.setdefault("_config_lock", RLock())
21
+ with lock:
22
+ depth = getattr(self, "_operation_depth", 0)
23
+ self._operation_depth = depth + 1
24
+ before = namespace_to_plain_dict(self.config)
25
+ previous_inputs = getattr(self, "_call_config_inputs", {})
26
+ bound = inspect.signature(method).bind(self, *args, **kwargs)
27
+ self._call_config_inputs = {**previous_inputs, **{
28
+ name.removesuffix("_set"): os.path.abspath(os.path.expanduser(str(value)))
29
+ for name, value in bound.arguments.items()
30
+ if name in {"train_set", "valid_set", "test_set"} and value is not None
31
+ }}
32
+ try:
33
+ if method.__name__ == "train" and self.is_remote:
34
+ raise ValueError("Remote training is not supported")
35
+ return method(self, *args, **kwargs)
36
+ finally:
37
+ self._operation_depth = depth
38
+ self._call_config_inputs = previous_inputs
39
+ if depth == 0:
40
+ # Local localization helpers write temporary input paths into
41
+ # config. Restore only those defaults, preserving model state.
42
+ current = namespace_to_plain_dict(self.config)
43
+ for split, cfg in _get(before, "DATA.common.splits", {}).items():
44
+ for key in ("annotation_path", "source_path"):
45
+ if key in cfg:
46
+ _set(current, f"DATA.common.splits.{split}.{key}", cfg[key])
47
+ self.config = dict_to_namespace(current)
48
+ self._refresh_config_references()
49
+ return wrapped
50
+
51
+
52
+ class ConfigurationMixin:
53
+ def _effective_config(self, config):
54
+ editor = getattr(self, "_config_editor", None)
55
+ result = dict_to_namespace(editor.apply_to(config)) if editor is not None else config
56
+ if getattr(self, "_call_config_inputs", None):
57
+ doc = namespace_to_plain_dict(result)
58
+ for split, path in self._call_config_inputs.items():
59
+ _set(doc, f"DATA.common.splits.{split}.annotation_path", path)
60
+ if split == "valid" and _get(doc, "DATA.common.splits.valid_data_frames"):
61
+ _set(doc, "DATA.common.splits.valid_data_frames.annotation_path", path)
62
+ result = dict_to_namespace(doc)
63
+ return result
64
+
65
+ def _refresh_config_references(self):
66
+ for owner in (getattr(self, "model", None), getattr(self, "trainer", None)):
67
+ if owner is None:
68
+ continue
69
+ for key in ("config", "cfg", "cfg_model"):
70
+ if hasattr(owner, key):
71
+ setattr(owner, key, self.config)
72
+
73
+ def get_config(self):
74
+ return deepcopy(namespace_to_plain_dict(self.config))
75
+
76
+ def config_options(self):
77
+ editor = deepcopy(getattr(self, "_config_editor", None)) or Config(self.get_config(), source=self.config_path)
78
+ # Inspection reflects normalization and runtime-selected checkpoint data.
79
+ editor._document = self.get_config()
80
+ return editor.options()
81
+
82
+ def update_config(self, **options):
83
+ lock = self.__dict__.setdefault("_config_lock", RLock())
84
+ if not lock.acquire(blocking=False):
85
+ raise RuntimeError("Cannot update configuration during an active operation")
86
+ try:
87
+ if getattr(self, "_operation_depth", 0):
88
+ raise RuntimeError("Cannot update configuration during an active operation")
89
+ original = getattr(self, "_config_editor", None)
90
+ editor = deepcopy(original) if original is not None else Config(self.get_config(), source=self.config_path)
91
+ # Carry current runtime values while retaining expressions and prior
92
+ # explicit updates from the editable source.
93
+ current = self.get_config()
94
+ editor._document = editor.apply_to(current)
95
+ if original is not None:
96
+ from opensportslib.core.config.editable import _leaves
97
+ for path, value in _leaves(original._document):
98
+ if isinstance(value, str) and "${" in value:
99
+ _set(editor._document, path, value)
100
+ previous = deepcopy(editor._updates)
101
+ editor.update(**options)
102
+ safe = []
103
+ for name, spec in editor._registry().items():
104
+ if not spec.requires_initialization and name.split(".")[0] in {"data", "training", "inference", "runtime", "scheduler", "training_sampling", "sft", "prompt"}:
105
+ safe.extend(spec.paths)
106
+ changed = [p for p, value in editor._updates.items() if p not in previous or previous[p] != value]
107
+ for path in changed:
108
+ if not any(path == p for p in safe):
109
+ raise ValueError(f"{path} requires a new model; update Config before initialization")
110
+ if self.is_remote:
111
+ editor.remote_overrides()
112
+ candidate = dict_to_namespace(editor.apply_to(current))
113
+ self.config = candidate
114
+ self._config_editor = editor
115
+ self._refresh_config_references()
116
+ return self
117
+ finally:
118
+ lock.release()
119
+
120
+ def _config_request_fields(self, fields):
121
+ fields = dict(fields)
122
+ editor = getattr(self, "_config_editor", None)
123
+ overrides = editor.remote_overrides() if editor is not None else {}
124
+ task_options = json.loads(fields.get("task_options") or "{}")
125
+ if "config_overrides" in task_options:
126
+ raise ValueError("Use Config.update(inference=...) for remote configuration overrides")
127
+ if not overrides:
128
+ return fields
129
+ query = urlencode({"model_id": fields.get("model_id", ""), "task_type": fields.get("task_type", "")})
130
+ try:
131
+ capabilities = self._open_request(Request(f"{self.remote}/config-capabilities?{query}"))
132
+ except Exception as exc:
133
+ raise ValueError("Server does not expose configuration capabilities; upgrade the server or omit overrides") from exc
134
+ if capabilities.get("version") != 1:
135
+ raise ValueError("Unsupported server configuration protocol version")
136
+ allowed = capabilities.get("options", {})
137
+ for name, value in overrides.items():
138
+ if name not in allowed:
139
+ raise ValueError(f"Server does not support inference.{name}")
140
+ spec = allowed[name]
141
+ if spec.get("minimum") is not None and value < spec["minimum"] or spec.get("maximum") is not None and value > spec["maximum"]:
142
+ raise ValueError(f"inference.{name} exceeds server limits")
143
+ task_options["config_overrides"] = {"version": 1, "inference": overrides}
144
+ fields["task_options"] = json.dumps(task_options)
145
+ return fields
@@ -4,6 +4,7 @@ import time
4
4
  from types import SimpleNamespace
5
5
 
6
6
  from opensportslib.apis.base_task_model import BaseTaskModel
7
+ from opensportslib.apis.configuration import config_operation
7
8
  from opensportslib.core.config.accessors import (
8
9
  get_data_classes,
9
10
  get_loader_backend,
@@ -290,6 +291,7 @@ class LocalizationModel(BaseTaskModel):
290
291
  "best_criterion_valid": best_criterion_valid,
291
292
  }
292
293
 
294
+ @config_operation
293
295
  def train(
294
296
  self,
295
297
  train_set=None,
@@ -327,7 +329,7 @@ class LocalizationModel(BaseTaskModel):
327
329
  # with explicit valid annotation overrides.
328
330
  self._set_split_path("valid_data_frames", valid_set)
329
331
 
330
- self.config = resolve_config_omega(self.config, weights=weights)
332
+ self.config = self._effective_config(resolve_config_omega(self.config, weights=weights))
331
333
  self.config = resolve_inference_class_metadata(self.config)
332
334
  effective_weights = weights if weights is not None else self.last_loaded_weights
333
335
  self._adapt_hf_backend_for_device(effective_weights)
@@ -418,14 +420,18 @@ class LocalizationModel(BaseTaskModel):
418
420
  logging.info(f"Total Execution Time is {time.time()-start} seconds")
419
421
  return self.best_checkpoint
420
422
 
423
+ @config_operation
421
424
  def infer(
422
425
  self,
423
426
  test_set=None,
424
427
  weights=None,
425
428
  use_wandb=True,
429
+ video_path: str | None = None,
426
430
  **kwargs,
427
431
  ):
428
432
  """Run model inference and return predictions in OSL JSON format."""
433
+ if test_set is not None and video_path is not None:
434
+ raise ValueError("Provide either `test_set` or `video_path`, not both.")
429
435
  remote_mode_provided = "remote_mode" in kwargs
430
436
  remote_mode = kwargs.pop("remote_mode", "full_test_set")
431
437
  if self.is_remote:
@@ -433,6 +439,17 @@ class LocalizationModel(BaseTaskModel):
433
439
  remote_task_options = kwargs.pop("remote_task_options", None)
434
440
  if kwargs:
435
441
  raise TypeError(f"Unsupported remote inference options: {', '.join(kwargs)}")
442
+ if video_path is not None:
443
+ if remote_mode != "full_test_set":
444
+ raise ValueError("`remote_mode=per_sample` requires `test_set`, not direct video input.")
445
+ job = self.submit_video_inference(
446
+ task_type="localization",
447
+ video_path=video_path,
448
+ model_id=remote_model_id,
449
+ task_options=remote_task_options,
450
+ )
451
+ self.last_remote_failures = []
452
+ return self.wait_for_remote_result(job["job_id"])["result"]["predictions"]
436
453
  test_set = self._resolve_split_path("test", test_set)
437
454
  if remote_mode == "per_sample":
438
455
  batch = self.submit_per_sample_inference(
@@ -458,6 +475,18 @@ class LocalizationModel(BaseTaskModel):
458
475
  return self.wait_for_remote_result(job["job_id"])["result"]["predictions"]
459
476
  if remote_mode_provided:
460
477
  raise ValueError("`remote_mode` is available only when `remote` is configured.")
478
+ if video_path is not None:
479
+ from opensportslib.core.utils.direct_video import direct_video_manifest
480
+ from opensportslib.core.utils.config import resolve_config_omega
481
+
482
+ manifest_config = self._effective_config(resolve_config_omega(self.config, weights=weights))
483
+ manifest_config = resolve_inference_class_metadata(manifest_config)
484
+ with direct_video_manifest(manifest_config, video_path, "localization") as manifest:
485
+ return self.infer(
486
+ test_set=manifest,
487
+ weights=weights,
488
+ use_wandb=use_wandb,
489
+ )
461
490
  from opensportslib.datasets.builder import build_dataset
462
491
  from opensportslib.models.builder import build_model
463
492
  from opensportslib.core.trainer.localization_trainer import build_inferer
@@ -473,7 +502,7 @@ class LocalizationModel(BaseTaskModel):
473
502
  test_set = self._resolve_split_path("test", test_set)
474
503
  self._set_split_path("test", test_set)
475
504
 
476
- self.config = resolve_config_omega(self.config, weights=weights)
505
+ self.config = self._effective_config(resolve_config_omega(self.config, weights=weights))
477
506
  self.config = resolve_inference_class_metadata(self.config)
478
507
  effective_weights = weights if weights is not None else self.last_loaded_weights
479
508
  self._adapt_hf_backend_for_device(effective_weights)
@@ -530,6 +559,7 @@ class LocalizationModel(BaseTaskModel):
530
559
  logging.info(f"Total Execution Time is {time.time()-start} seconds")
531
560
  return predictions
532
561
 
562
+ @config_operation
533
563
  def evaluate(
534
564
  self,
535
565
  test_set=None,
@@ -551,7 +581,7 @@ class LocalizationModel(BaseTaskModel):
551
581
 
552
582
  test_set = self._resolve_split_path("test", test_set)
553
583
  self._set_split_path("test", test_set)
554
- self.config = resolve_config_omega(self.config, weights=weights)
584
+ self.config = self._effective_config(resolve_config_omega(self.config, weights=weights))
555
585
  self.config = resolve_inference_class_metadata(self.config)
556
586
  effective_weights = weights if weights is not None else self.last_loaded_weights
557
587
  self._adapt_hf_backend_for_device(effective_weights)
opensportslib/apis/vqa.py CHANGED
@@ -8,6 +8,7 @@ import os
8
8
  from typing import Any
9
9
 
10
10
  from opensportslib.apis.base_task_model import BaseTaskModel
11
+ from opensportslib.apis.configuration import config_operation
11
12
  from opensportslib.core.config.accessors import get_split_annotation_path, get_system_gpu_count, get_train_execution, get_vqa_backend
12
13
  from opensportslib.core.utils.config import expand, resolve_config_omega
13
14
 
@@ -132,6 +133,7 @@ class VQAModel(BaseTaskModel):
132
133
  self.last_loaded_weights = weights
133
134
  self.best_checkpoint = weights
134
135
 
136
+ @config_operation
135
137
  def train(
136
138
  self,
137
139
  train_set: str | None = None,
@@ -144,7 +146,7 @@ class VQAModel(BaseTaskModel):
144
146
  ) -> str | None:
145
147
  del kwargs
146
148
 
147
- self.config = resolve_config_omega(self.config, weights=weights)
149
+ self.config = self._effective_config(resolve_config_omega(self.config, weights=weights))
148
150
  execution = get_train_execution(self.config)
149
151
  backend = str(execution.get("training_backend", "placeholder")).lower()
150
152
  vqa_backend = get_vqa_backend(self.config)
@@ -227,6 +229,7 @@ class VQAModel(BaseTaskModel):
227
229
  "Only 'xvars_videochatgpt_lora' and 'qwen_xvars_lora' are supported."
228
230
  )
229
231
 
232
+ @config_operation
230
233
  def infer(
231
234
  self,
232
235
  test_set: str | None = None,
@@ -309,7 +312,7 @@ class VQAModel(BaseTaskModel):
309
312
  if direct_requested and (not video_path or not str(question or "").strip()):
310
313
  raise ValueError("Direct VQA inference requires both `video_path` and a non-empty `question`.")
311
314
 
312
- self.config = resolve_config_omega(self.config, weights=weights)
315
+ self.config = self._effective_config(resolve_config_omega(self.config, weights=weights))
313
316
  backend = get_vqa_backend(self.config)
314
317
  effective_weights = weights if weights is not None else self.last_loaded_weights
315
318
  _set_model_checkpoint_path(self.config, effective_weights)
@@ -346,6 +349,7 @@ class VQAModel(BaseTaskModel):
346
349
  self._init_wandb(use_wandb=use_wandb)
347
350
  return self.trainer.infer(model, test_data, use_wandb=use_wandb)
348
351
 
352
+ @config_operation
349
353
  def evaluate(
350
354
  self,
351
355
  test_set: str | None = None,
@@ -358,7 +362,7 @@ class VQAModel(BaseTaskModel):
358
362
  from opensportslib.core.trainer.vqa_trainer import Trainer_VQA
359
363
  from opensportslib.datasets.builder import build_dataset
360
364
 
361
- self.config = resolve_config_omega(self.config, weights=weights)
365
+ self.config = self._effective_config(resolve_config_omega(self.config, weights=weights))
362
366
  test_set = self._resolve_split_path("test", test_set)
363
367
  test_data = build_dataset(self.config, test_set, None, split="test")
364
368
  self._init_wandb(use_wandb=use_wandb)