feat: Auto install requirements

This commit is contained in:
Vinay-Umrethe
2026-06-04 19:29:49 +05:30
parent b5d9d72849
commit f4be1abd04
+81 -4
View File
@@ -3,9 +3,69 @@
# ruff: noqa: E402
import importlib
import json
import subprocess
import sys
import tempfile
from pathlib import Path
from typing import Any
from urllib.request import urlopen
from .config import Settings
_BOOTSTRAP_REPRODUCTION_INFORMATION: dict[str, Any] | None = None
def _get_reproduce_path() -> str | None:
if "--reproduce" in sys.argv:
index = sys.argv.index("--reproduce")
if index + 1 < len(sys.argv):
return sys.argv[index + 1]
for arg in sys.argv[1:]:
if arg.startswith("--reproduce="):
return arg.split("=", 1)[1]
return None
def _load_reproduction_information_bootstrap(path: str) -> dict[str, Any]:
if path.lower().startswith(("http://", "https://")):
path = path.replace("/blob/", "/raw/")
path = path.replace("/src/branch/", "/raw/branch/")
with urlopen(path, timeout=30) as response:
return json.loads(response.read().decode("utf-8"))
else:
return json.loads(Path(path).read_text(encoding="utf-8"))
def _install_reproduction_requirements(requirements: dict[str, str]):
with tempfile.NamedTemporaryFile(
"w",
suffix=".txt",
delete=False,
encoding="utf-8",
) as requirements_file:
for package, package_version in sorted(requirements.items()):
requirements_file.write(f"{package}=={package_version}\n")
requirements_path = Path(requirements_file.name)
try:
subprocess.check_call(
[
sys.executable,
"-m",
"pip",
"install",
"-q",
"-r",
str(requirements_path),
]
)
finally:
requirements_path.unlink(missing_ok=True)
importlib.invalidate_caches()
def _is_help_invocation() -> bool:
@@ -13,6 +73,21 @@ def _is_help_invocation() -> bool:
return "-h" in args or "--help" in args
if not _is_help_invocation():
reproduce_path = _get_reproduce_path()
if reproduce_path:
_BOOTSTRAP_REPRODUCTION_INFORMATION = _load_reproduction_information_bootstrap(
reproduce_path
)
requirements = _BOOTSTRAP_REPRODUCTION_INFORMATION.get("environment", {}).get(
"requirements"
)
if requirements:
_install_reproduction_requirements(requirements)
from .config import Settings
# Parse and handle CLI help before importing heavyweight ML/runtime dependencies.
if _is_help_invocation():
Settings() # ty:ignore[missing-argument]
@@ -36,8 +111,6 @@ import warnings
from dataclasses import asdict
from importlib.metadata import version
from os.path import commonprefix
from pathlib import Path
from typing import Any
import huggingface_hub
import lm_eval
@@ -222,7 +295,11 @@ def run():
if settings.reproduce is not None:
print(f"Loading reproduction information from [bold]{settings.reproduce}[/]...")
# FIXME: "Reproduction"/"reproducibility" name inconsistency!
reproduction_information = load_reproduction_information(settings.reproduce)
reproduction_information = (
_BOOTSTRAP_REPRODUCTION_INFORMATION
if _BOOTSTRAP_REPRODUCTION_INFORMATION is not None
else load_reproduction_information(settings.reproduce)
)
if reproduction_information["version"] not in ["1"]:
print(