opensportslib 0.3.0.dev15__tar.gz → 0.3.0.dev17__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 (195) hide show
  1. {opensportslib-0.3.0.dev15/opensportslib.egg-info → opensportslib-0.3.0.dev17}/PKG-INFO +1 -1
  2. opensportslib-0.3.0.dev17/opensportslib/setup/setup.py +291 -0
  3. {opensportslib-0.3.0.dev15 → opensportslib-0.3.0.dev17/opensportslib.egg-info}/PKG-INFO +1 -1
  4. {opensportslib-0.3.0.dev15 → opensportslib-0.3.0.dev17}/pyproject.toml +1 -1
  5. {opensportslib-0.3.0.dev15 → opensportslib-0.3.0.dev17}/tests/test_setup_cli.py +31 -0
  6. opensportslib-0.3.0.dev15/opensportslib/setup/setup.py +0 -211
  7. {opensportslib-0.3.0.dev15 → opensportslib-0.3.0.dev17}/LICENSE +0 -0
  8. {opensportslib-0.3.0.dev15 → opensportslib-0.3.0.dev17}/LICENSE-COMMERCIAL +0 -0
  9. {opensportslib-0.3.0.dev15 → opensportslib-0.3.0.dev17}/MANIFEST.in +0 -0
  10. {opensportslib-0.3.0.dev15 → opensportslib-0.3.0.dev17}/README.md +0 -0
  11. {opensportslib-0.3.0.dev15 → opensportslib-0.3.0.dev17}/examples/quickstart/basic_classification.py +0 -0
  12. {opensportslib-0.3.0.dev15 → opensportslib-0.3.0.dev17}/examples/quickstart/basic_localization.py +0 -0
  13. {opensportslib-0.3.0.dev15 → opensportslib-0.3.0.dev17}/examples/quickstart/basic_vqa.py +0 -0
  14. {opensportslib-0.3.0.dev15 → opensportslib-0.3.0.dev17}/opensportslib/__init__.py +0 -0
  15. {opensportslib-0.3.0.dev15 → opensportslib-0.3.0.dev17}/opensportslib/adaptation/__init__.py +0 -0
  16. {opensportslib-0.3.0.dev15 → opensportslib-0.3.0.dev17}/opensportslib/adaptation/spotta.py +0 -0
  17. {opensportslib-0.3.0.dev15 → opensportslib-0.3.0.dev17}/opensportslib/apis/__init__.py +0 -0
  18. {opensportslib-0.3.0.dev15 → opensportslib-0.3.0.dev17}/opensportslib/apis/base_task_model.py +0 -0
  19. {opensportslib-0.3.0.dev15 → opensportslib-0.3.0.dev17}/opensportslib/apis/classification.py +0 -0
  20. {opensportslib-0.3.0.dev15 → opensportslib-0.3.0.dev17}/opensportslib/apis/localization.py +0 -0
  21. {opensportslib-0.3.0.dev15 → opensportslib-0.3.0.dev17}/opensportslib/apis/vqa.py +0 -0
  22. {opensportslib-0.3.0.dev15 → opensportslib-0.3.0.dev17}/opensportslib/cli.py +0 -0
  23. {opensportslib-0.3.0.dev15 → opensportslib-0.3.0.dev17}/opensportslib/configs/classification/default.yaml +0 -0
  24. {opensportslib-0.3.0.dev15 → opensportslib-0.3.0.dev17}/opensportslib/configs/classification/sngar_frames.yaml +0 -0
  25. {opensportslib-0.3.0.dev15 → opensportslib-0.3.0.dev17}/opensportslib/configs/classification/sngar_tracking.yaml +0 -0
  26. {opensportslib-0.3.0.dev15 → opensportslib-0.3.0.dev17}/opensportslib/configs/classification/video.yaml +0 -0
  27. {opensportslib-0.3.0.dev15 → opensportslib-0.3.0.dev17}/opensportslib/configs/default.yaml +0 -0
  28. {opensportslib-0.3.0.dev15 → opensportslib-0.3.0.dev17}/opensportslib/configs/localization/calf_resnetpca512.yaml +0 -0
  29. {opensportslib-0.3.0.dev15 → opensportslib-0.3.0.dev17}/opensportslib/configs/localization/default.yaml +0 -0
  30. {opensportslib-0.3.0.dev15 → opensportslib-0.3.0.dev17}/opensportslib/configs/localization/e2e_spotta.yaml +0 -0
  31. {opensportslib-0.3.0.dev15 → opensportslib-0.3.0.dev17}/opensportslib/configs/localization/h5_header_distance.yaml +0 -0
  32. {opensportslib-0.3.0.dev15 → opensportslib-0.3.0.dev17}/opensportslib/configs/localization/h5_header_skeleton.yaml +0 -0
  33. {opensportslib-0.3.0.dev15 → opensportslib-0.3.0.dev17}/opensportslib/configs/localization/netvladpp_resnetpca512.yaml +0 -0
  34. {opensportslib-0.3.0.dev15 → opensportslib-0.3.0.dev17}/opensportslib/configs/localization/tracking_action_spotting.yaml +0 -0
  35. {opensportslib-0.3.0.dev15 → opensportslib-0.3.0.dev17}/opensportslib/configs/localization/video_dali.yaml +0 -0
  36. {opensportslib-0.3.0.dev15 → opensportslib-0.3.0.dev17}/opensportslib/configs/localization/video_ocv.yaml +0 -0
  37. {opensportslib-0.3.0.dev15 → opensportslib-0.3.0.dev17}/opensportslib/configs/vqa/default.yaml +0 -0
  38. {opensportslib-0.3.0.dev15 → opensportslib-0.3.0.dev17}/opensportslib/configs/vqa/qwen.yaml +0 -0
  39. {opensportslib-0.3.0.dev15 → opensportslib-0.3.0.dev17}/opensportslib/configs/vqa/qwen3_vl_native.yaml +0 -0
  40. {opensportslib-0.3.0.dev15 → opensportslib-0.3.0.dev17}/opensportslib/configs/vqa/qwen_lora.yaml +0 -0
  41. {opensportslib-0.3.0.dev15 → opensportslib-0.3.0.dev17}/opensportslib/configs/vqa/qwen_sngar_frames.yaml +0 -0
  42. {opensportslib-0.3.0.dev15 → opensportslib-0.3.0.dev17}/opensportslib/configs/vqa/xvars.yaml +0 -0
  43. {opensportslib-0.3.0.dev15 → opensportslib-0.3.0.dev17}/opensportslib/core/__init__.py +0 -0
  44. {opensportslib-0.3.0.dev15 → opensportslib-0.3.0.dev17}/opensportslib/core/config/__init__.py +0 -0
  45. {opensportslib-0.3.0.dev15 → opensportslib-0.3.0.dev17}/opensportslib/core/config/accessors.py +0 -0
  46. {opensportslib-0.3.0.dev15 → opensportslib-0.3.0.dev17}/opensportslib/core/config/conflicts.py +0 -0
  47. {opensportslib-0.3.0.dev15 → opensportslib-0.3.0.dev17}/opensportslib/core/config/loader.py +0 -0
  48. {opensportslib-0.3.0.dev15 → opensportslib-0.3.0.dev17}/opensportslib/core/config/migrate.py +0 -0
  49. {opensportslib-0.3.0.dev15 → opensportslib-0.3.0.dev17}/opensportslib/core/config/migrations/__init__.py +0 -0
  50. {opensportslib-0.3.0.dev15 → opensportslib-0.3.0.dev17}/opensportslib/core/config/migrations/legacy_to_canonical.py +0 -0
  51. {opensportslib-0.3.0.dev15 → opensportslib-0.3.0.dev17}/opensportslib/core/config/runtime_adapter.py +0 -0
  52. {opensportslib-0.3.0.dev15 → opensportslib-0.3.0.dev17}/opensportslib/core/config/schema.py +0 -0
  53. {opensportslib-0.3.0.dev15 → opensportslib-0.3.0.dev17}/opensportslib/core/config/schemas/__init__.py +0 -0
  54. {opensportslib-0.3.0.dev15 → opensportslib-0.3.0.dev17}/opensportslib/core/config/schemas/schema_canonical.py +0 -0
  55. {opensportslib-0.3.0.dev15 → opensportslib-0.3.0.dev17}/opensportslib/core/config/schemas/schema_legacy.py +0 -0
  56. {opensportslib-0.3.0.dev15 → opensportslib-0.3.0.dev17}/opensportslib/core/config/validate.py +0 -0
  57. {opensportslib-0.3.0.dev15 → opensportslib-0.3.0.dev17}/opensportslib/core/loss/__init__.py +0 -0
  58. {opensportslib-0.3.0.dev15 → opensportslib-0.3.0.dev17}/opensportslib/core/loss/builder.py +0 -0
  59. {opensportslib-0.3.0.dev15 → opensportslib-0.3.0.dev17}/opensportslib/core/loss/calf.py +0 -0
  60. {opensportslib-0.3.0.dev15 → opensportslib-0.3.0.dev17}/opensportslib/core/loss/ce.py +0 -0
  61. {opensportslib-0.3.0.dev15 → opensportslib-0.3.0.dev17}/opensportslib/core/loss/combine.py +0 -0
  62. {opensportslib-0.3.0.dev15 → opensportslib-0.3.0.dev17}/opensportslib/core/loss/nll.py +0 -0
  63. {opensportslib-0.3.0.dev15 → opensportslib-0.3.0.dev17}/opensportslib/core/optimizer/__init__.py +0 -0
  64. {opensportslib-0.3.0.dev15 → opensportslib-0.3.0.dev17}/opensportslib/core/optimizer/builder.py +0 -0
  65. {opensportslib-0.3.0.dev15 → opensportslib-0.3.0.dev17}/opensportslib/core/sampler/weighted_sampler.py +0 -0
  66. {opensportslib-0.3.0.dev15 → opensportslib-0.3.0.dev17}/opensportslib/core/scheduler/__init__.py +0 -0
  67. {opensportslib-0.3.0.dev15 → opensportslib-0.3.0.dev17}/opensportslib/core/scheduler/builder.py +0 -0
  68. {opensportslib-0.3.0.dev15 → opensportslib-0.3.0.dev17}/opensportslib/core/trainer/__init__.py +0 -0
  69. {opensportslib-0.3.0.dev15 → opensportslib-0.3.0.dev17}/opensportslib/core/trainer/classification_trainer.py +0 -0
  70. {opensportslib-0.3.0.dev15 → opensportslib-0.3.0.dev17}/opensportslib/core/trainer/localization_trainer.py +0 -0
  71. {opensportslib-0.3.0.dev15 → opensportslib-0.3.0.dev17}/opensportslib/core/trainer/vqa_trainer.py +0 -0
  72. {opensportslib-0.3.0.dev15 → opensportslib-0.3.0.dev17}/opensportslib/core/utils/checkpoint.py +0 -0
  73. {opensportslib-0.3.0.dev15 → opensportslib-0.3.0.dev17}/opensportslib/core/utils/config.py +0 -0
  74. {opensportslib-0.3.0.dev15 → opensportslib-0.3.0.dev17}/opensportslib/core/utils/config_normalize.py +0 -0
  75. {opensportslib-0.3.0.dev15 → opensportslib-0.3.0.dev17}/opensportslib/core/utils/data.py +0 -0
  76. {opensportslib-0.3.0.dev15 → opensportslib-0.3.0.dev17}/opensportslib/core/utils/ddp.py +0 -0
  77. {opensportslib-0.3.0.dev15 → opensportslib-0.3.0.dev17}/opensportslib/core/utils/default_args.py +0 -0
  78. {opensportslib-0.3.0.dev15 → opensportslib-0.3.0.dev17}/opensportslib/core/utils/hf_runtime.py +0 -0
  79. {opensportslib-0.3.0.dev15 → opensportslib-0.3.0.dev17}/opensportslib/core/utils/lightning.py +0 -0
  80. {opensportslib-0.3.0.dev15 → opensportslib-0.3.0.dev17}/opensportslib/core/utils/load_annotations.py +0 -0
  81. {opensportslib-0.3.0.dev15 → opensportslib-0.3.0.dev17}/opensportslib/core/utils/seed.py +0 -0
  82. {opensportslib-0.3.0.dev15 → opensportslib-0.3.0.dev17}/opensportslib/core/utils/video_processing.py +0 -0
  83. {opensportslib-0.3.0.dev15 → opensportslib-0.3.0.dev17}/opensportslib/core/utils/wandb.py +0 -0
  84. {opensportslib-0.3.0.dev15 → opensportslib-0.3.0.dev17}/opensportslib/datasets/__init__.py +0 -0
  85. {opensportslib-0.3.0.dev15 → opensportslib-0.3.0.dev17}/opensportslib/datasets/builder.py +0 -0
  86. {opensportslib-0.3.0.dev15 → opensportslib-0.3.0.dev17}/opensportslib/datasets/classification_dataset.py +0 -0
  87. {opensportslib-0.3.0.dev15 → opensportslib-0.3.0.dev17}/opensportslib/datasets/localization_dataset.py +0 -0
  88. {opensportslib-0.3.0.dev15 → opensportslib-0.3.0.dev17}/opensportslib/datasets/utils/__init__.py +0 -0
  89. {opensportslib-0.3.0.dev15 → opensportslib-0.3.0.dev17}/opensportslib/datasets/utils/h5_tracking.py +0 -0
  90. {opensportslib-0.3.0.dev15 → opensportslib-0.3.0.dev17}/opensportslib/datasets/utils/tracking.py +0 -0
  91. {opensportslib-0.3.0.dev15 → opensportslib-0.3.0.dev17}/opensportslib/datasets/vqa_dataset.py +0 -0
  92. {opensportslib-0.3.0.dev15 → opensportslib-0.3.0.dev17}/opensportslib/legacy_config/classification.yaml +0 -0
  93. {opensportslib-0.3.0.dev15 → opensportslib-0.3.0.dev17}/opensportslib/legacy_config/localization-e2e-ocv.yaml +0 -0
  94. {opensportslib-0.3.0.dev15 → opensportslib-0.3.0.dev17}/opensportslib/legacy_config/localization-json_calf_resnetpca512.yaml +0 -0
  95. {opensportslib-0.3.0.dev15 → opensportslib-0.3.0.dev17}/opensportslib/legacy_config/localization-json_netvlad++_resnetpca512.yaml +0 -0
  96. {opensportslib-0.3.0.dev15 → opensportslib-0.3.0.dev17}/opensportslib/legacy_config/localization.yaml +0 -0
  97. {opensportslib-0.3.0.dev15 → opensportslib-0.3.0.dev17}/opensportslib/legacy_config/sngar-frames.yaml +0 -0
  98. {opensportslib-0.3.0.dev15 → opensportslib-0.3.0.dev17}/opensportslib/legacy_config/sngar-tracking.yaml +0 -0
  99. {opensportslib-0.3.0.dev15 → opensportslib-0.3.0.dev17}/opensportslib/metrics/classification_metric.py +0 -0
  100. {opensportslib-0.3.0.dev15 → opensportslib-0.3.0.dev17}/opensportslib/metrics/localization_metric.py +0 -0
  101. {opensportslib-0.3.0.dev15 → opensportslib-0.3.0.dev17}/opensportslib/metrics/vqa_metric.py +0 -0
  102. {opensportslib-0.3.0.dev15 → opensportslib-0.3.0.dev17}/opensportslib/models/__init__.py +0 -0
  103. {opensportslib-0.3.0.dev15 → opensportslib-0.3.0.dev17}/opensportslib/models/backbones/builder.py +0 -0
  104. {opensportslib-0.3.0.dev15 → opensportslib-0.3.0.dev17}/opensportslib/models/base/contextaware.py +0 -0
  105. {opensportslib-0.3.0.dev15 → opensportslib-0.3.0.dev17}/opensportslib/models/base/e2e.py +0 -0
  106. {opensportslib-0.3.0.dev15 → opensportslib-0.3.0.dev17}/opensportslib/models/base/learnablepooling.py +0 -0
  107. {opensportslib-0.3.0.dev15 → opensportslib-0.3.0.dev17}/opensportslib/models/base/qwen_vl_native.py +0 -0
  108. {opensportslib-0.3.0.dev15 → opensportslib-0.3.0.dev17}/opensportslib/models/base/qwen_xvars.py +0 -0
  109. {opensportslib-0.3.0.dev15 → opensportslib-0.3.0.dev17}/opensportslib/models/base/rule_based.py +0 -0
  110. {opensportslib-0.3.0.dev15 → opensportslib-0.3.0.dev17}/opensportslib/models/base/tracking.py +0 -0
  111. {opensportslib-0.3.0.dev15 → opensportslib-0.3.0.dev17}/opensportslib/models/base/vars.py +0 -0
  112. {opensportslib-0.3.0.dev15 → opensportslib-0.3.0.dev17}/opensportslib/models/base/video.py +0 -0
  113. {opensportslib-0.3.0.dev15 → opensportslib-0.3.0.dev17}/opensportslib/models/base/video_chatgpt_compat.py +0 -0
  114. {opensportslib-0.3.0.dev15 → opensportslib-0.3.0.dev17}/opensportslib/models/base/video_mae.py +0 -0
  115. {opensportslib-0.3.0.dev15 → opensportslib-0.3.0.dev17}/opensportslib/models/base/xvars_videochatgpt.py +0 -0
  116. {opensportslib-0.3.0.dev15 → opensportslib-0.3.0.dev17}/opensportslib/models/builder.py +0 -0
  117. {opensportslib-0.3.0.dev15 → opensportslib-0.3.0.dev17}/opensportslib/models/heads/builder.py +0 -0
  118. {opensportslib-0.3.0.dev15 → opensportslib-0.3.0.dev17}/opensportslib/models/neck/builder.py +0 -0
  119. {opensportslib-0.3.0.dev15 → opensportslib-0.3.0.dev17}/opensportslib/models/utils/common.py +0 -0
  120. {opensportslib-0.3.0.dev15 → opensportslib-0.3.0.dev17}/opensportslib/models/utils/impl/__init__.py +0 -0
  121. {opensportslib-0.3.0.dev15 → opensportslib-0.3.0.dev17}/opensportslib/models/utils/impl/asformer.py +0 -0
  122. {opensportslib-0.3.0.dev15 → opensportslib-0.3.0.dev17}/opensportslib/models/utils/impl/calf.py +0 -0
  123. {opensportslib-0.3.0.dev15 → opensportslib-0.3.0.dev17}/opensportslib/models/utils/impl/gsm.py +0 -0
  124. {opensportslib-0.3.0.dev15 → opensportslib-0.3.0.dev17}/opensportslib/models/utils/impl/gtad.py +0 -0
  125. {opensportslib-0.3.0.dev15 → opensportslib-0.3.0.dev17}/opensportslib/models/utils/impl/tsm.py +0 -0
  126. {opensportslib-0.3.0.dev15 → opensportslib-0.3.0.dev17}/opensportslib/models/utils/litebase.py +0 -0
  127. {opensportslib-0.3.0.dev15 → opensportslib-0.3.0.dev17}/opensportslib/models/utils/modules.py +0 -0
  128. {opensportslib-0.3.0.dev15 → opensportslib-0.3.0.dev17}/opensportslib/models/utils/shift.py +0 -0
  129. {opensportslib-0.3.0.dev15 → opensportslib-0.3.0.dev17}/opensportslib/models/utils/utils.py +0 -0
  130. {opensportslib-0.3.0.dev15 → opensportslib-0.3.0.dev17}/opensportslib/models/utils/vqa_prediction_priors.py +0 -0
  131. {opensportslib-0.3.0.dev15 → opensportslib-0.3.0.dev17}/opensportslib/models/utils/vqa_prompting.py +0 -0
  132. {opensportslib-0.3.0.dev15 → opensportslib-0.3.0.dev17}/opensportslib/models/utils/xvars_clip_index.py +0 -0
  133. {opensportslib-0.3.0.dev15 → opensportslib-0.3.0.dev17}/opensportslib/tools/__init__.py +0 -0
  134. {opensportslib-0.3.0.dev15 → opensportslib-0.3.0.dev17}/opensportslib/tools/_common.py +0 -0
  135. {opensportslib-0.3.0.dev15 → opensportslib-0.3.0.dev17}/opensportslib/tools/hf_transfer.py +0 -0
  136. {opensportslib-0.3.0.dev15 → opensportslib-0.3.0.dev17}/opensportslib/tools/osl_json_to_parquet.py +0 -0
  137. {opensportslib-0.3.0.dev15 → opensportslib-0.3.0.dev17}/opensportslib/tools/parquet_to_osl_json.py +0 -0
  138. {opensportslib-0.3.0.dev15 → opensportslib-0.3.0.dev17}/opensportslib.egg-info/SOURCES.txt +0 -0
  139. {opensportslib-0.3.0.dev15 → opensportslib-0.3.0.dev17}/opensportslib.egg-info/dependency_links.txt +0 -0
  140. {opensportslib-0.3.0.dev15 → opensportslib-0.3.0.dev17}/opensportslib.egg-info/entry_points.txt +0 -0
  141. {opensportslib-0.3.0.dev15 → opensportslib-0.3.0.dev17}/opensportslib.egg-info/requires.txt +0 -0
  142. {opensportslib-0.3.0.dev15 → opensportslib-0.3.0.dev17}/opensportslib.egg-info/top_level.txt +0 -0
  143. {opensportslib-0.3.0.dev15 → opensportslib-0.3.0.dev17}/scripts/run_h5_header_rule_inference.py +0 -0
  144. {opensportslib-0.3.0.dev15 → opensportslib-0.3.0.dev17}/setup.cfg +0 -0
  145. {opensportslib-0.3.0.dev15 → opensportslib-0.3.0.dev17}/tests/conftest.py +0 -0
  146. {opensportslib-0.3.0.dev15 → opensportslib-0.3.0.dev17}/tests/release/__init__.py +0 -0
  147. {opensportslib-0.3.0.dev15 → opensportslib-0.3.0.dev17}/tests/release/_release_common.py +0 -0
  148. {opensportslib-0.3.0.dev15 → opensportslib-0.3.0.dev17}/tests/release/test_classification_release.py +0 -0
  149. {opensportslib-0.3.0.dev15 → opensportslib-0.3.0.dev17}/tests/release/test_localization_release.py +0 -0
  150. {opensportslib-0.3.0.dev15 → opensportslib-0.3.0.dev17}/tests/release/test_vqa_release.py +0 -0
  151. {opensportslib-0.3.0.dev15 → opensportslib-0.3.0.dev17}/tests/test_classification_dataset_paths.py +0 -0
  152. {opensportslib-0.3.0.dev15 → opensportslib-0.3.0.dev17}/tests/test_classification_trainer_dataloader.py +0 -0
  153. {opensportslib-0.3.0.dev15 → opensportslib-0.3.0.dev17}/tests/test_config_architecture.py +0 -0
  154. {opensportslib-0.3.0.dev15 → opensportslib-0.3.0.dev17}/tests/test_config_split_override_sync.py +0 -0
  155. {opensportslib-0.3.0.dev15 → opensportslib-0.3.0.dev17}/tests/test_config_utils_smoke.py +0 -0
  156. {opensportslib-0.3.0.dev15 → opensportslib-0.3.0.dev17}/tests/test_conversion_tools.py +0 -0
  157. {opensportslib-0.3.0.dev15 → opensportslib-0.3.0.dev17}/tests/test_extract_xvars_features.py +0 -0
  158. {opensportslib-0.3.0.dev15 → opensportslib-0.3.0.dev17}/tests/test_h5_header_rule_spotter.py +0 -0
  159. {opensportslib-0.3.0.dev15 → opensportslib-0.3.0.dev17}/tests/test_h5_header_skeleton_spotter.py +0 -0
  160. {opensportslib-0.3.0.dev15 → opensportslib-0.3.0.dev17}/tests/test_h5_tracking_dataset.py +0 -0
  161. {opensportslib-0.3.0.dev15 → opensportslib-0.3.0.dev17}/tests/test_hf_transfer_tools.py +0 -0
  162. {opensportslib-0.3.0.dev15 → opensportslib-0.3.0.dev17}/tests/test_localization_dali_filenames.py +0 -0
  163. {opensportslib-0.3.0.dev15 → opensportslib-0.3.0.dev17}/tests/test_localization_hf_backend_override.py +0 -0
  164. {opensportslib-0.3.0.dev15 → opensportslib-0.3.0.dev17}/tests/test_localization_intervals.py +0 -0
  165. {opensportslib-0.3.0.dev15 → opensportslib-0.3.0.dev17}/tests/test_package_smoke.py +0 -0
  166. {opensportslib-0.3.0.dev15 → opensportslib-0.3.0.dev17}/tests/test_pretrained_config_merge_policy.py +0 -0
  167. {opensportslib-0.3.0.dev15 → opensportslib-0.3.0.dev17}/tests/test_public_apis_smoke.py +0 -0
  168. {opensportslib-0.3.0.dev15 → opensportslib-0.3.0.dev17}/tests/test_spotta_e2e.py +0 -0
  169. {opensportslib-0.3.0.dev15 → opensportslib-0.3.0.dev17}/tests/test_subset_train_infer_integration.py +0 -0
  170. {opensportslib-0.3.0.dev15 → opensportslib-0.3.0.dev17}/tests/test_task_model_api_contract.py +0 -0
  171. {opensportslib-0.3.0.dev15 → opensportslib-0.3.0.dev17}/tests/test_vqa_api.py +0 -0
  172. {opensportslib-0.3.0.dev15 → opensportslib-0.3.0.dev17}/tests/test_vqa_metrics_semantic.py +0 -0
  173. {opensportslib-0.3.0.dev15 → opensportslib-0.3.0.dev17}/tests/test_vqa_qwen_xvars.py +0 -0
  174. {opensportslib-0.3.0.dev15 → opensportslib-0.3.0.dev17}/tests/test_vqa_training_lora.py +0 -0
  175. {opensportslib-0.3.0.dev15 → opensportslib-0.3.0.dev17}/tests/test_vqa_xvars_videochatgpt.py +0 -0
  176. {opensportslib-0.3.0.dev15 → opensportslib-0.3.0.dev17}/tools/convert/build_sn_vqa_2026_vqa.py +0 -0
  177. {opensportslib-0.3.0.dev15 → opensportslib-0.3.0.dev17}/tools/convert/build_sngar_spotting.py +0 -0
  178. {opensportslib-0.3.0.dev15 → opensportslib-0.3.0.dev17}/tools/convert/build_soccernet_gar.py +0 -0
  179. {opensportslib-0.3.0.dev15 → opensportslib-0.3.0.dev17}/tools/convert/build_soccernet_gar_action_spotting.py +0 -0
  180. {opensportslib-0.3.0.dev15 → opensportslib-0.3.0.dev17}/tools/convert/build_soccernet_gar_vqa.py +0 -0
  181. {opensportslib-0.3.0.dev15 → opensportslib-0.3.0.dev17}/tools/convert/build_xvars_indexes.py +0 -0
  182. {opensportslib-0.3.0.dev15 → opensportslib-0.3.0.dev17}/tools/convert/extract_xvars_clip_features.py +0 -0
  183. {opensportslib-0.3.0.dev15 → opensportslib-0.3.0.dev17}/tools/convert/osl_json_to_parquet_webdataset.py +0 -0
  184. {opensportslib-0.3.0.dev15 → opensportslib-0.3.0.dev17}/tools/convert/parquet_webdataset_to_osl_json.py +0 -0
  185. {opensportslib-0.3.0.dev15 → opensportslib-0.3.0.dev17}/tools/convert/sngar_dataset_card.py +0 -0
  186. {opensportslib-0.3.0.dev15 → opensportslib-0.3.0.dev17}/tools/convert/sngar_events.py +0 -0
  187. {opensportslib-0.3.0.dev15 → opensportslib-0.3.0.dev17}/tools/convert/verify_sngar_spotting.py +0 -0
  188. {opensportslib-0.3.0.dev15 → opensportslib-0.3.0.dev17}/tools/download/download_hf_repo.py +0 -0
  189. {opensportslib-0.3.0.dev15 → opensportslib-0.3.0.dev17}/tools/download/download_osl_hf.py +0 -0
  190. {opensportslib-0.3.0.dev15 → opensportslib-0.3.0.dev17}/tools/download/push_sngar_spotting.py +0 -0
  191. {opensportslib-0.3.0.dev15 → opensportslib-0.3.0.dev17}/tools/download/upload_osl_hf.py +0 -0
  192. {opensportslib-0.3.0.dev15 → opensportslib-0.3.0.dev17}/tools/training/classification.py +0 -0
  193. {opensportslib-0.3.0.dev15 → opensportslib-0.3.0.dev17}/tools/training/localization.py +0 -0
  194. {opensportslib-0.3.0.dev15 → opensportslib-0.3.0.dev17}/tools/training/vqa.py +0 -0
  195. {opensportslib-0.3.0.dev15 → opensportslib-0.3.0.dev17}/tools/upload/upload_model_hf.py +0 -0
@@ -1,6 +1,6 @@
1
1
  Metadata-Version: 2.4
2
2
  Name: opensportslib
3
- Version: 0.3.0.dev15
3
+ Version: 0.3.0.dev17
4
4
  Summary: OpenSportsLib is the professional library, designed for advanced video understanding in sports. It provides state-of-the-art tools for action recognition, spotting, retrieval, and captioning, making it ideal for researchers, analysts, and developers working with sports video data.
5
5
  Author: Jeet Vora
6
6
  Requires-Python: >=3.12
@@ -0,0 +1,291 @@
1
+ import platform
2
+ import subprocess
3
+ import sys
4
+
5
+
6
+ CUDA_WHEEL_VERSIONS = {
7
+ "cu126": (12, 6),
8
+ "cu128": (12, 8),
9
+ "cu130": (13, 0),
10
+ }
11
+ MIN_SUPPORTED_COMPUTE_CAPABILITY = (5, 0)
12
+ LEGACY_GPU_MAX_COMPUTE_CAPABILITY = (7, 4)
13
+ LEGACY_GPU_CUDA_WHEEL = "cu126"
14
+ LEGACY_GPU_CUDA_WHEEL_MAX_COMPUTE_CAPABILITY = (9, 0)
15
+ CUDA13_REQUIRED_MIN_COMPUTE_CAPABILITY = (10, 0)
16
+ LEGACY_GPU_TORCH_PACKAGES = (
17
+ "torch==2.10.0",
18
+ "torchvision==0.25.0",
19
+ "torchaudio==2.10.0",
20
+ )
21
+
22
+ XVARS_DEPENDENCY_PINS = {
23
+ "transformers": "4.38.2",
24
+ "peft": "0.9.0",
25
+ "tokenizers": "0.15.2",
26
+ "accelerate": "0.27.2",
27
+ "trl": "0.10.1",
28
+ }
29
+
30
+ QWEN_DEPENDENCY_PINS = {
31
+ "transformers": "5.13.0",
32
+ "peft": "0.19.0",
33
+ "tokenizers": "0.22.1",
34
+ "accelerate": "1.14.0",
35
+ "trl": "1.7.1",
36
+ }
37
+
38
+ def get_cuda_version():
39
+ try:
40
+ output = subprocess.check_output(["nvidia-smi"]).decode()
41
+
42
+ for line in output.split("\n"):
43
+ if "CUDA Version" in line:
44
+ ver = line.split("CUDA Version:")[1].strip().split()[0]
45
+ print(f"CUDA Version found : {ver}")
46
+ cuda_tag = f"cu{ver.replace('.', '')}"
47
+ return ver, cuda_tag
48
+ except Exception:
49
+ return None, None
50
+ return None, None
51
+
52
+
53
+ def get_gpu_compute_capabilities():
54
+ try:
55
+ output = subprocess.check_output(
56
+ ["nvidia-smi", "--query-gpu=compute_cap", "--format=csv,noheader"],
57
+ text=True,
58
+ )
59
+ except Exception:
60
+ return []
61
+
62
+ capabilities = []
63
+ for line in output.splitlines():
64
+ value = line.strip()
65
+ if not value:
66
+ continue
67
+ try:
68
+ major, minor = value.split(".", 1)
69
+ capabilities.append((int(major), int(minor)))
70
+ except ValueError:
71
+ print(f"Ignoring unrecognized GPU compute capability: {value}")
72
+ return capabilities
73
+
74
+
75
+ def select_cuda_wheel(cuda_version, compute_capabilities):
76
+ if not cuda_version:
77
+ return "cpu"
78
+
79
+ try:
80
+ driver_version = tuple(int(part) for part in cuda_version.split(".", 1))
81
+ except ValueError as exc:
82
+ raise RuntimeError(f"Unable to parse CUDA version reported by nvidia-smi: {cuda_version}") from exc
83
+
84
+ if any(capability < MIN_SUPPORTED_COMPUTE_CAPABILITY for capability in compute_capabilities):
85
+ raise RuntimeError(
86
+ "OpenSportsLib PyTorch wheels require GPUs with compute capability "
87
+ f"{MIN_SUPPORTED_COMPUTE_CAPABILITY[0]}.{MIN_SUPPORTED_COMPUTE_CAPABILITY[1]} or newer. "
88
+ f"Detected: {compute_capabilities}."
89
+ )
90
+
91
+ compatible_tags = [
92
+ tag for tag, wheel_version in CUDA_WHEEL_VERSIONS.items() if wheel_version <= driver_version
93
+ ]
94
+ has_legacy_gpu = compute_capabilities and any(
95
+ capability <= LEGACY_GPU_MAX_COMPUTE_CAPABILITY for capability in compute_capabilities
96
+ )
97
+ if has_legacy_gpu:
98
+ if platform.machine().lower() in {"aarch64", "arm64"}:
99
+ raise RuntimeError(
100
+ "Official Linux ARM64 CUDA 12.6 PyTorch wheels support Ampere and newer GPUs only. "
101
+ "Use an x86_64 host or build PyTorch from source for the visible legacy GPU architecture."
102
+ )
103
+ if any(
104
+ capability > LEGACY_GPU_CUDA_WHEEL_MAX_COMPUTE_CAPABILITY
105
+ for capability in compute_capabilities
106
+ ):
107
+ raise RuntimeError(
108
+ "Visible GPUs require incompatible prebuilt PyTorch wheels. Set CUDA_VISIBLE_DEVICES to "
109
+ "either the legacy GPUs or the newer GPUs, then run setup again."
110
+ )
111
+ compatible_tags = [tag for tag in compatible_tags if tag == LEGACY_GPU_CUDA_WHEEL]
112
+
113
+ if any(
114
+ capability >= CUDA13_REQUIRED_MIN_COMPUTE_CAPABILITY
115
+ for capability in compute_capabilities
116
+ ):
117
+ compatible_tags = [tag for tag in compatible_tags if tag == "cu130"]
118
+
119
+ if not compatible_tags:
120
+ if any(
121
+ capability >= CUDA13_REQUIRED_MIN_COMPUTE_CAPABILITY
122
+ for capability in compute_capabilities
123
+ ):
124
+ raise RuntimeError(
125
+ "The visible GPU architecture requires the CUDA 13.0 PyTorch wheel. "
126
+ "Update the NVIDIA driver so nvidia-smi reports CUDA 13.0 or newer, "
127
+ "then run setup again."
128
+ )
129
+ raise RuntimeError(
130
+ f"CUDA {cuda_version} is too old for the supported PyTorch wheels: "
131
+ f"{', '.join(CUDA_WHEEL_VERSIONS)}."
132
+ )
133
+
134
+ selected = max(compatible_tags, key=lambda tag: CUDA_WHEEL_VERSIONS[tag])
135
+ print(
136
+ "Selected PyTorch wheel "
137
+ f"{selected} for CUDA {cuda_version} and GPU compute capabilities {compute_capabilities or 'unknown'}"
138
+ )
139
+ return selected
140
+
141
+
142
+ def select_torch_packages(compute_capabilities):
143
+ if compute_capabilities and any(
144
+ capability <= LEGACY_GPU_MAX_COMPUTE_CAPABILITY for capability in compute_capabilities
145
+ ):
146
+ return LEGACY_GPU_TORCH_PACKAGES
147
+ return ("torch", "torchvision", "torchaudio")
148
+
149
+
150
+ CUDA_VERSION, _DETECTED_CUDA_TAG = get_cuda_version()
151
+ GPU_COMPUTE_CAPABILITIES = get_gpu_compute_capabilities()
152
+ CUDA_TAG = select_cuda_wheel(CUDA_VERSION, GPU_COMPUTE_CAPABILITIES)
153
+
154
+
155
+ def install_xvars_dependencies(DEPENDENCY_PINS):
156
+ python = sys.executable
157
+ packages = list(DEPENDENCY_PINS)
158
+ pinned_packages = [f"{name}=={version}" for name, version in DEPENDENCY_PINS.items()]
159
+
160
+ print(f"\nInstalling {list(DEPENDENCY_PINS.keys())} dependency overrides...\n")
161
+ print("This overrides the default Hugging Face dependency set with XVars-compatible versions.")
162
+ subprocess.call([python, "-m", "pip", "uninstall", "-y", *packages])
163
+ subprocess.check_call([python, "-m", "pip", "install", *pinned_packages])
164
+ print("Dependencies installed successfully.")
165
+
166
+ def install_torch():
167
+ python = sys.executable
168
+ subprocess.call([python, "-m", "pip", "uninstall", "-y", "torch", "torchvision", "torchaudio"])
169
+ packages = select_torch_packages(GPU_COMPUTE_CAPABILITIES)
170
+
171
+ subprocess.check_call([
172
+ python, "-m", "pip", "install",
173
+ *packages,
174
+ "--index-url",
175
+ f"https://download.pytorch.org/whl/{CUDA_TAG}",
176
+ ])
177
+ print(f"\nSuccess with {CUDA_TAG}: {', '.join(packages)}")
178
+ return CUDA_TAG
179
+
180
+ def install_dali():
181
+
182
+ python = sys.executable
183
+
184
+ print("\nInstalling dali extras...\n")
185
+
186
+ # DALI (only if GPU)
187
+ if CUDA_VERSION:
188
+
189
+ if CUDA_TAG == "cu130":
190
+ subprocess.check_call([
191
+ python, "-m", "pip", "install",
192
+ "nvidia-dali-cuda130"
193
+ ])
194
+
195
+ # CuPy (CUDA-aware but auto-resolves internally)
196
+ subprocess.check_call([
197
+ python, "-m", "pip", "install",
198
+ "cupy-cuda130"
199
+ ])
200
+ else:
201
+ subprocess.check_call([
202
+ python, "-m", "pip", "install",
203
+ "nvidia-dali-cuda120"
204
+ ])
205
+
206
+ # CuPy (CUDA-aware but auto-resolves internally)
207
+ subprocess.check_call([
208
+ python, "-m", "pip", "install",
209
+ "cupy-cuda12x"
210
+ ])
211
+
212
+ def install_pyg():
213
+ import torch
214
+ from packaging import version
215
+
216
+ python = sys.executable
217
+ torch_version = "2.10.0" if version.parse(torch.__version__.split("+")[0]) > version.parse("2.10.0") else torch.__version__.split("+")[0]
218
+ cuda_tag = CUDA_TAG
219
+ print("\nInstalling Py-Geometric ecosystem...\n")
220
+ if cuda_tag == "cpu":
221
+ url = f"https://data.pyg.org/whl/torch-{torch_version}+cpu.html"
222
+ else:
223
+ url = f"https://data.pyg.org/whl/torch-{torch_version}+{cuda_tag}.html"
224
+
225
+ subprocess.check_call([
226
+ python, "-m", "pip", "install",
227
+ "torch-geometric", "-f", url
228
+ ])
229
+ subprocess.check_call([
230
+ python, "-m", "pip", "install",
231
+ "torch-scatter", "-f", url
232
+ ])
233
+ subprocess.check_call([
234
+ python, "-m", "pip", "install",
235
+ "torch-sparse", "-f", url
236
+ ])
237
+ subprocess.check_call([
238
+ python, "-m", "pip", "install",
239
+ "torch-cluster", "-f", url
240
+ ])
241
+ subprocess.check_call([
242
+ python, "-m", "pip", "install",
243
+ "torch-spline-conv", "-f", url
244
+ ])
245
+
246
+ def install_extras(dali=False, pyg=False):
247
+ if dali:
248
+ install_dali()
249
+ print("NVIDIA DALI installed successfully.")
250
+ if pyg:
251
+ install_pyg()
252
+ print("PyTorch Geometric installed successfully.")
253
+
254
+
255
+ def verify():
256
+ import torch
257
+
258
+ print("\n Verifying installation...\n")
259
+ print("Torch:", torch.__version__)
260
+
261
+ if torch.cuda.is_available():
262
+ print("CUDA available")
263
+ print("GPU:", torch.cuda.get_device_name(0))
264
+ else:
265
+ print("Running on CPU")
266
+
267
+ def setup(dali=False, pyg=False, vqa_xvars=False, vqa_qwen=False):
268
+ install_torch()
269
+ install_extras(dali=dali, pyg=pyg)
270
+ if vqa_xvars:
271
+ install_xvars_dependencies(XVARS_DEPENDENCY_PINS)
272
+ if vqa_qwen:
273
+ install_xvars_dependencies(QWEN_DEPENDENCY_PINS)
274
+ verify()
275
+
276
+
277
+ # ----------------------------
278
+ # CLI entry
279
+ # ----------------------------
280
+ if __name__ == "__main__":
281
+ import argparse
282
+
283
+ parser = argparse.ArgumentParser()
284
+ parser.add_argument("--dali", action="store_true")
285
+ parser.add_argument("--pyg", action="store_true")
286
+ parser.add_argument("--vqa_xvars", action="store_true")
287
+ parser.add_argument("--vqa_qwen", action="store_true")
288
+
289
+ args = parser.parse_args()
290
+
291
+ setup(dali=args.dali, pyg=args.pyg, vqa_xvars=args.vqa_xvars, vqa_qwen=args.vqa_qwen)
@@ -1,6 +1,6 @@
1
1
  Metadata-Version: 2.4
2
2
  Name: opensportslib
3
- Version: 0.3.0.dev15
3
+ Version: 0.3.0.dev17
4
4
  Summary: OpenSportsLib is the professional library, designed for advanced video understanding in sports. It provides state-of-the-art tools for action recognition, spotting, retrieval, and captioning, making it ideal for researchers, analysts, and developers working with sports video data.
5
5
  Author: Jeet Vora
6
6
  Requires-Python: >=3.12
@@ -4,7 +4,7 @@ build-backend = "setuptools.build_meta"
4
4
 
5
5
  [project]
6
6
  name = "opensportslib"
7
- version = "0.3.0.dev15"
7
+ version = "0.3.0.dev17"
8
8
  description = "OpenSportsLib is the professional library, designed for advanced video understanding in sports. It provides state-of-the-art tools for action recognition, spotting, retrieval, and captioning, making it ideal for researchers, analysts, and developers working with sports video data."
9
9
  readme = "README.md"
10
10
  requires-python = ">=3.12"
@@ -1,9 +1,40 @@
1
1
  from __future__ import annotations
2
2
 
3
+ import pytest
4
+
3
5
  from opensportslib import cli
4
6
  from opensportslib.setup import setup as setup_lib
5
7
 
6
8
 
9
+ def test_select_cuda_wheel_uses_cu126_for_pre_sm75_gpu_with_cuda_13():
10
+ assert setup_lib.select_cuda_wheel("13.0", [(7, 0)]) == "cu126"
11
+
12
+
13
+ def test_select_cuda_wheel_uses_cu126_for_pascal_with_cuda_13():
14
+ assert setup_lib.select_cuda_wheel("13.0", [(6, 0)]) == "cu126"
15
+
16
+
17
+ def test_select_torch_packages_pins_pre_sm75_gpu_compatibility_stack():
18
+ assert setup_lib.select_torch_packages([(7, 0)]) == (
19
+ "torch==2.10.0",
20
+ "torchvision==0.25.0",
21
+ "torchaudio==2.10.0",
22
+ )
23
+
24
+
25
+ def test_select_cuda_wheel_uses_cu130_for_dgx_spark():
26
+ assert setup_lib.select_cuda_wheel("13.0", [(12, 1)]) == "cu130"
27
+
28
+
29
+ def test_select_cuda_wheel_rejects_new_architecture_without_cuda_13_driver():
30
+ with pytest.raises(RuntimeError, match="CUDA 13.0"):
31
+ setup_lib.select_cuda_wheel("12.8", [(10, 0)])
32
+
33
+
34
+ def test_select_cuda_wheel_uses_highest_driver_compatible_wheel():
35
+ assert setup_lib.select_cuda_wheel("12.8", [(8, 0)]) == "cu128"
36
+
37
+
7
38
  def test_cli_setup_forwards_xvars_flag(monkeypatch):
8
39
  captured: dict[str, object] = {}
9
40
 
@@ -1,211 +0,0 @@
1
- import subprocess
2
- import sys
3
-
4
-
5
- CUDA_SUPPORT = [
6
- "cu126",
7
- "cu128",
8
- "cu130",
9
- "cpu"
10
- ]
11
-
12
- XVARS_DEPENDENCY_PINS = {
13
- "transformers": "4.38.2",
14
- "peft": "0.9.0",
15
- "tokenizers": "0.15.2",
16
- "accelerate": "0.27.2",
17
- "trl": "0.10.1",
18
- }
19
-
20
- QWEN_DEPENDENCY_PINS = {
21
- "transformers": "5.13.0",
22
- "peft": "0.19.0",
23
- "tokenizers": "0.22.1",
24
- "accelerate": "1.14.0",
25
- "trl": "1.7.1",
26
- }
27
-
28
- def get_cuda_version():
29
- try:
30
- output = subprocess.check_output(["nvidia-smi"]).decode()
31
-
32
- for line in output.split("\n"):
33
- if "CUDA Version" in line:
34
- ver = line.split("CUDA Version:")[1].strip().split()[0]
35
- print(f"CUDA Version found : {ver}")
36
- cuda_tag = f"cu{ver.replace('.', '')}"
37
- return ver, cuda_tag
38
- except Exception:
39
- return None, None
40
-
41
-
42
- CUDA_VERSION, CUDA_TAG = get_cuda_version()
43
-
44
- def get_cpu_tag():
45
- if not CUDA_VERSION:
46
- return "cpu"
47
-
48
-
49
- def install_xvars_dependencies(DEPENDENCY_PINS):
50
- python = sys.executable
51
- packages = list(DEPENDENCY_PINS)
52
- pinned_packages = [f"{name}=={version}" for name, version in DEPENDENCY_PINS.items()]
53
-
54
- print(f"\nInstalling {list(DEPENDENCY_PINS.keys())} dependency overrides...\n")
55
- print("This overrides the default Hugging Face dependency set with XVars-compatible versions.")
56
- subprocess.call([python, "-m", "pip", "uninstall", "-y", *packages])
57
- subprocess.check_call([python, "-m", "pip", "install", *pinned_packages])
58
- print("Dependencies installed successfully.")
59
-
60
- def install_torch():
61
- python = sys.executable
62
- subprocess.call([python, "-m", "pip", "uninstall", "-y", "torch", "torchvision"])
63
-
64
- if CUDA_TAG == "cu130":
65
- cuda = "cu130"
66
- subprocess.check_call([
67
- python, "-m", "pip", "install",
68
- "torch", "torchvision", "torchaudio",
69
- "--index-url",
70
- f"https://download.pytorch.org/whl/{cuda}"
71
- ])
72
- print(f"\nSuccess with {cuda}")
73
- return cuda
74
- for cuda in CUDA_SUPPORT:
75
-
76
- print(f"\n Trying installation: {cuda}\n")
77
- try:
78
- if get_cpu_tag() == "cpu":
79
- subprocess.check_call([
80
- python, "-m", "pip", "install",
81
- "torch", "torchvision", "torchaudio"
82
- ])
83
- else:
84
-
85
- subprocess.check_call([
86
- python, "-m", "pip", "install",
87
- "torch", "torchvision", "torchaudio",
88
- "--index-url",
89
- f"https://download.pytorch.org/whl/{cuda}"
90
- ])
91
- print(f"\nSuccess with {cuda}")
92
- return cuda
93
-
94
- except Exception as e:
95
- print(f"Failed with {cuda}: {e}")
96
-
97
- raise RuntimeError("All CUDA installs failed")
98
-
99
- def install_dali():
100
-
101
- python = sys.executable
102
-
103
- print("\nInstalling dali extras...\n")
104
-
105
- # DALI (only if GPU)
106
- if CUDA_VERSION:
107
-
108
- if CUDA_TAG == "cu130":
109
- subprocess.check_call([
110
- python, "-m", "pip", "install",
111
- "nvidia-dali-cuda130"
112
- ])
113
-
114
- # CuPy (CUDA-aware but auto-resolves internally)
115
- subprocess.check_call([
116
- python, "-m", "pip", "install",
117
- "cupy-cuda130"
118
- ])
119
- else:
120
- subprocess.check_call([
121
- python, "-m", "pip", "install",
122
- "nvidia-dali-cuda120"
123
- ])
124
-
125
- # CuPy (CUDA-aware but auto-resolves internally)
126
- subprocess.check_call([
127
- python, "-m", "pip", "install",
128
- "cupy-cuda12x"
129
- ])
130
-
131
- def install_pyg():
132
- import torch
133
- from packaging import version
134
-
135
- python = sys.executable
136
- torch_version = "2.10.0" if version.parse(torch.__version__.split("+")[0]) > version.parse("2.10.0") else torch.__version__.split("+")[0]
137
- cuda_tag = next((f"cu{CUDA_VERSION.replace('.', '')}" for _ in [0] if CUDA_VERSION), CUDA_SUPPORT[0])
138
- cuda_tag = cuda_tag if cuda_tag in CUDA_SUPPORT else CUDA_SUPPORT[0]
139
- print("\nInstalling Py-Geometric ecosystem...\n")
140
- if cuda_tag == "cpu":
141
- url = f"https://data.pyg.org/whl/torch-{torch_version}+cpu.html"
142
- else:
143
- url = f"https://data.pyg.org/whl/torch-{torch_version}+{cuda_tag}.html"
144
-
145
- subprocess.check_call([
146
- python, "-m", "pip", "install",
147
- "torch-geometric", "-f", url
148
- ])
149
- subprocess.check_call([
150
- python, "-m", "pip", "install",
151
- "torch-scatter", "-f", url
152
- ])
153
- subprocess.check_call([
154
- python, "-m", "pip", "install",
155
- "torch-sparse", "-f", url
156
- ])
157
- subprocess.check_call([
158
- python, "-m", "pip", "install",
159
- "torch-cluster", "-f", url
160
- ])
161
- subprocess.check_call([
162
- python, "-m", "pip", "install",
163
- "torch-spline-conv", "-f", url
164
- ])
165
-
166
- def install_extras(dali=False, pyg=False):
167
- if dali:
168
- install_dali()
169
- print("NVIDIA DALI installed successfully.")
170
- if pyg:
171
- install_pyg()
172
- print("PyTorch Geometric installed successfully.")
173
-
174
-
175
- def verify():
176
- import torch
177
-
178
- print("\n Verifying installation...\n")
179
- print("Torch:", torch.__version__)
180
-
181
- if torch.cuda.is_available():
182
- print("CUDA available")
183
- print("GPU:", torch.cuda.get_device_name(0))
184
- else:
185
- print("Running on CPU")
186
-
187
- def setup(dali=False, pyg=False, vqa_xvars=False, vqa_qwen=False):
188
- install_torch()
189
- install_extras(dali=dali, pyg=pyg)
190
- if vqa_xvars:
191
- install_xvars_dependencies(XVARS_DEPENDENCY_PINS)
192
- if vqa_qwen:
193
- install_xvars_dependencies(QWEN_DEPENDENCY_PINS)
194
- verify()
195
-
196
-
197
- # ----------------------------
198
- # CLI entry
199
- # ----------------------------
200
- if __name__ == "__main__":
201
- import argparse
202
-
203
- parser = argparse.ArgumentParser()
204
- parser.add_argument("--dali", action="store_true")
205
- parser.add_argument("--pyg", action="store_true")
206
- parser.add_argument("--vqa_xvars", action="store_true")
207
- parser.add_argument("--vqa_qwen", action="store_true")
208
-
209
- args = parser.parse_args()
210
-
211
- setup(dali=args.dali, pyg=args.pyg, vqa_xvars=args.vqa_xvars, vqa_qwen=args.vqa_qwen)