vaghawan's picture
Add multi-speaker 3-speaker epoch-10 inference bundle
5827d9c verified
Raw History Blame Contribute Delete
5.55 kB
#!/usr/bin/env python3
"""Load KEY=VALUE settings from config.env (project root).
CLI flags always win over config.env values.
"""
from __future__ import annotations
import os
import shlex
from pathlib import Path
ROOT = Path(__file__).resolve().parent
DEFAULT_ENV_PATH = ROOT / "config.env"
# Always take these from config.env so a broken shell mirror/token cannot win.
_CONFIG_FORCE_KEYS = frozenset(
{
"HF_ENDPOINT",
"HF_TOKEN",
"HUGGINGFACE_ACCESS_TOKEN",
"HUGGING_FACE_HUB_TOKEN",
"HUGGINGFACE_TOKEN",
"HF_HUB_DOWNLOAD_TIMEOUT",
"HF_HUB_ETAG_TIMEOUT",
"HF_DOWNLOAD_RETRIES",
}
)
def load_config_env(path: Path | None = None, *, override: bool = False) -> Path | None:
"""Parse config.env into os.environ.
Returns the path loaded, or None if missing.
Does not override existing environment variables unless override=True,
except for `_CONFIG_FORCE_KEYS` (HF endpoint/token/timeouts).
"""
env_path = (path or DEFAULT_ENV_PATH).resolve()
if not env_path.is_file():
return None
for raw in env_path.read_text(encoding="utf-8").splitlines():
line = raw.strip()
if not line or line.startswith("#"):
continue
if line.startswith("export "):
line = line[len("export ") :].strip()
if "=" not in line:
continue
key, value = line.split("=", 1)
key = key.strip()
value = value.strip()
if not key:
continue
if len(value) >= 2 and value[0] == value[-1] and value[0] in "\"'":
value = value[1:-1]
# Skip blank assignments (e.g. FOO=) so they don't poison Hub clients
if value == "":
continue
if override or key not in os.environ or key in _CONFIG_FORCE_KEYS:
os.environ[key] = value
return env_path
def env_str(name: str, default: str | None = None) -> str | None:
v = os.getenv(name)
if v is None or v == "":
return default
return v
def env_int(name: str, default: int) -> int:
v = os.getenv(name)
if v is None or v == "":
return default
return int(v)
def env_float(name: str, default: float) -> float:
v = os.getenv(name)
if v is None or v == "":
return default
return float(v)
def env_path(name: str, default: str | Path) -> Path:
v = os.getenv(name)
return Path(v) if v else Path(default)
def env_list(name: str, default: list[str] | None = None) -> list[str]:
"""Split a config value into a list (whitespace or comma separated)."""
v = os.getenv(name)
if v is None or v.strip() == "":
return list(default or [])
normalized = v.replace(",", " ")
parts = shlex.split(normalized)
return [p for p in parts if p]
def ensure_config_loaded(path: Path | None = None) -> None:
loaded = load_config_env(path)
if loaded:
print(f"[config] loaded {loaded}")
else:
example = ROOT / "config.env.example"
print(
f"[config] no config.env found"
+ (f" (copy from {example.name})" if example.is_file() else "")
)
apply_hf_token()
def _sync_hf_hub_endpoint(endpoint: str) -> None:
"""Keep huggingface_hub / datasets in sync if they were already imported."""
endpoint = endpoint.rstrip("/")
try:
import huggingface_hub.constants as hub_constants
hub_constants.ENDPOINT = endpoint
hub_constants.HUGGINGFACE_CO_URL_TEMPLATE = endpoint + "/{repo_id}/resolve/{revision}/{filename}"
except Exception:
pass
try:
import datasets.config as ds_config
ds_config.HF_ENDPOINT = endpoint
ds_config.HUB_DATASETS_URL = endpoint + "/datasets/{repo_id}/resolve/{revision}/{path}"
except Exception:
pass
def apply_hf_token() -> None:
"""Normalize HF token aliases and sync Hub endpoint.
HF_ENDPOINT from config.env already wins over the shell via
`_CONFIG_FORCE_KEYS` — no need to `unset HF_ENDPOINT` before runs.
"""
endpoint = (os.environ.get("HF_ENDPOINT") or "").strip().rstrip("/")
if endpoint and not endpoint.startswith(("http://", "https://")):
del os.environ["HF_ENDPOINT"]
endpoint = ""
if not endpoint:
endpoint = "https://huggingface.co"
os.environ["HF_ENDPOINT"] = endpoint
_sync_hf_hub_endpoint(endpoint)
print(f"[config] HF_ENDPOINT={endpoint}")
# Longer timeouts help IncompleteRead on slow/unstable links
for key, default in (
("HF_HUB_DOWNLOAD_TIMEOUT", "120"),
("HF_HUB_ETAG_TIMEOUT", "60"),
("HF_DOWNLOAD_RETRIES", "8"),
):
val = env_str(key, default)
if val:
os.environ[key] = val
token = (
env_str("HF_TOKEN")
or env_str("HUGGINGFACE_ACCESS_TOKEN")
or env_str("HUGGING_FACE_HUB_TOKEN")
or env_str("HUGGINGFACE_TOKEN")
)
if not token:
print("[config] HF token not set (optional for public repos)")
return
os.environ["HF_TOKEN"] = token
os.environ["HUGGING_FACE_HUB_TOKEN"] = token
if not os.getenv("HUGGINGFACE_ACCESS_TOKEN"):
os.environ["HUGGINGFACE_ACCESS_TOKEN"] = token
print(f"[config] HF token loaded ({token[:6]}…{token[-4:] if len(token) > 10 else '****'})")
try:
from huggingface_hub import login
login(token=token, add_to_git_credential=False)
except Exception as e:
print(f"[config] huggingface_hub.login skipped: {e}")