Team Ai
Apppublic

evalstate/diffusers-pr-api

sourceHugging Faceupdated 6mo agoView on Hugging Face
0likes
save_cache.py116 linesDownload Raw Back to app
1from __future__ import annotations2 3from collections.abc import Callable4from pathlib import Path5from typing import Any, Protocol, cast6 7from huggingface_hub import HfApi8 9from slop_farmer.config import SaveCacheOptions10from slop_farmer.data.parquet_io import read_json11from slop_farmer.data.snapshot_paths import ROOT_MANIFEST_FILENAME, resolve_snapshot_dir_from_output12 13ANALYSIS_STATE_DIRNAME = "analysis-state"14 15 16class HubApiLike(Protocol):17    def create_repo(18        self,19        repo_id: str,20        *,21        repo_type: str,22        private: bool,23        exist_ok: bool,24    ) -> None: ...25 26    def upload_folder(27        self,28        *,29        repo_id: str,30        folder_path: Path,31        path_in_repo: str,32        repo_type: str,33        commit_message: str,34    ) -> None: ...35 36 37def run_save_cache(options: SaveCacheOptions) -> dict[str, Any]:38    snapshot_dir = resolve_snapshot_dir_from_output(options.output_dir, options.snapshot_dir)39    return save_analysis_cache(40        snapshot_dir=snapshot_dir,41        hf_repo_id=options.hf_repo_id,42        private=options.private_hf_repo,43    )44 45 46def save_analysis_cache(47    *,48    snapshot_dir: Path,49    hf_repo_id: str,50    private: bool,51    log: Callable[[str], None] | None = None,52) -> dict[str, Any]:53    return _save_analysis_cache_api(54        cast("HubApiLike", HfApi()),55        snapshot_dir=snapshot_dir,56        hf_repo_id=hf_repo_id,57        private=private,58        log=log,59    )60 61 62def _save_analysis_cache_api(63    api: HubApiLike,64    *,65    snapshot_dir: Path,66    hf_repo_id: str,67    private: bool,68    log: Callable[[str], None] | None = None,69) -> dict[str, Any]:70    cache_dir = snapshot_dir / ANALYSIS_STATE_DIRNAME71    if not cache_dir.exists():72        raise FileNotFoundError(f"Analysis cache directory is missing: {cache_dir}")73    if not cache_dir.is_dir():74        raise NotADirectoryError(f"Analysis cache path is not a directory: {cache_dir}")75    artifact_paths = _cache_artifact_paths(cache_dir)76    if not artifact_paths:77        raise ValueError(f"Analysis cache directory is empty: {cache_dir}")78 79    manifest_path = snapshot_dir / ROOT_MANIFEST_FILENAME80    manifest = read_json(manifest_path) if manifest_path.exists() else {}81    if not isinstance(manifest, dict):82        raise ValueError(f"Snapshot manifest at {manifest_path} must contain a JSON object.")83    snapshot_id = str(manifest.get("snapshot_id") or snapshot_dir.name).strip()84    repo = str(manifest.get("repo") or "").strip()85 86    if log:87        log(f"Ensuring Hub dataset repo exists: {hf_repo_id}")88    api.create_repo(hf_repo_id, repo_type="dataset", private=private, exist_ok=True)89    if log:90        log(f"Saving analysis cache for snapshot {snapshot_id}")91    api.upload_folder(92        repo_id=hf_repo_id,93        folder_path=cache_dir,94        path_in_repo=ANALYSIS_STATE_DIRNAME,95        repo_type="dataset",96        commit_message=f"Save analysis cache for snapshot {snapshot_id}",97    )98    result = {99        "dataset_id": hf_repo_id,100        "snapshot_id": snapshot_id,101        "artifact_paths": [f"{ANALYSIS_STATE_DIRNAME}/{path}" for path in artifact_paths],102    }103    if repo:104        result["repo"] = repo105    if log:106        log(f"Saved analysis cache to {hf_repo_id}")107    return result108 109 110def _cache_artifact_paths(cache_dir: Path) -> list[str]:111    return sorted(112        str(path.relative_to(cache_dir).as_posix())113        for path in cache_dir.rglob("*")114        if path.is_file()115    )116