File size: 3,996 Bytes
932bc69 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 | # ============================================================================
# AUTO-GENERATED — do not edit directly.
# Source of truth: src/speculators/provenance.py
# Regenerate with: make style (or: python scripts/sync_provenance.py)
# ============================================================================
"""Shared provenance helpers for reproducibility artifacts.
Used by training, evaluation, and vLLM-launch scripts to record
command lines, git state, and package versions.
"""
from __future__ import annotations
import importlib.metadata
import importlib.util
import os
import subprocess
import tempfile
from pathlib import Path
TRACKED_PACKAGES = (
"speculators",
"vllm",
"transformers",
"torch",
"compressed-tensors",
)
def atomic_write(path: Path, content: str) -> None:
"""Write *content* to *path* atomically via tempfile + rename."""
fd, tmp = tempfile.mkstemp(dir=path.parent, prefix=f".{path.name}_", suffix=".tmp")
tmp_path = Path(tmp)
try:
with os.fdopen(fd, "w") as f:
f.write(content)
tmp_path.replace(path)
finally:
if tmp_path.exists():
tmp_path.unlink()
def pkg_version(name: str) -> str:
"""Return the installed version of *name*, or ``'not installed'``."""
try:
return importlib.metadata.version(name)
except importlib.metadata.PackageNotFoundError:
return "not installed"
def find_repo_root(start: Path) -> Path | None:
"""Walk up from *start* to the nearest directory containing ``.git``."""
try:
d = start.resolve()
if d.is_file():
d = d.parent
while d != d.parent:
if (d / ".git").exists():
return d
d = d.parent
except OSError:
pass
return None
def find_package_repo(package_name: str) -> Path | None:
"""Find the git repo root for an installed editable package."""
try:
spec = importlib.util.find_spec(package_name)
if spec and spec.origin:
return find_repo_root(Path(spec.origin))
except (ModuleNotFoundError, ValueError):
pass
return None
def git_sha(repo_root: Path | None) -> str:
"""Return the HEAD SHA of *repo_root*, or ``'unknown'`` on failure."""
if repo_root is None:
return "unknown"
try:
result = subprocess.run(
["git", "rev-parse", "HEAD"], # noqa: S607
capture_output=True,
text=True,
cwd=repo_root,
timeout=5,
check=False,
)
return result.stdout.strip() or "unknown"
except (OSError, subprocess.TimeoutExpired):
return "unknown"
def git_diff(repo_root: Path | None, *, timeout: int = 30) -> str:
"""Return ``git diff HEAD`` for *repo_root*, or empty string on failure."""
if repo_root is None:
return ""
try:
result = subprocess.run(
["git", "diff", "HEAD"], # noqa: S607
capture_output=True,
text=True,
cwd=repo_root,
timeout=timeout,
check=False,
)
return result.stdout.strip() if result.returncode == 0 else ""
except (OSError, subprocess.TimeoutExpired):
return ""
def run_git(args: list[str], cwd: str | Path, *, timeout: int = 5) -> str:
"""Run a git command and return stdout, or empty string on failure."""
try:
result = subprocess.run( # noqa: S603
args,
capture_output=True,
text=True,
cwd=str(cwd),
timeout=timeout,
check=False,
)
return result.stdout.strip() if result.returncode == 0 else ""
except (OSError, subprocess.TimeoutExpired):
return ""
def package_versions(packages: tuple[str, ...] = TRACKED_PACKAGES) -> list[str]:
"""Return ``['# pkg: version', ...]`` header lines for *packages*."""
return [f"# {p}: {pkg_version(p)}" for p in packages]
|