evalstate/diffusers-pr-api
0
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 