Team Ai
Apppublic

evalstate/diffusers-pr-api

sourceHugging Faceupdated 6mo agoView on Hugging Face
0likes
snapshot_materialize.py540 linesDownload Raw Back to data
1from __future__ import annotations2 3import json4import shutil5import urllib.parse6import urllib.request7from datetime import UTC, datetime8from pathlib import Path, PurePosixPath9from typing import Any10 11from huggingface_hub import HfApi, hf_hub_download12 13from slop_farmer.data.http import urlopen_with_retry14from slop_farmer.data.parquet_io import read_json, write_text15from slop_farmer.data.snapshot_paths import (16    CONTRIBUTOR_ARTIFACT_FILENAMES,17    CURRENT_ANALYSIS_MANIFEST_PATH,18    LEGACY_ANALYSIS_FILENAMES,19    PR_SCOPE_CLUSTERS_FILENAME,20    RAW_TABLE_FILENAMES,21    README_FILENAME,22    ROOT_MANIFEST_FILENAME,23    SNAPSHOTS_LATEST_PATH,24    STATE_WATERMARK_PATH,25    load_archived_analysis_run_manifest,26    load_current_analysis_manifest,27    repo_relative_path_to_local,28)29 30 31def materialize_hf_dataset_snapshot(32    *,33    repo_id: str,34    local_dir: Path,35    revision: str | None = None,36) -> Path:37    info = _hf_dataset_info(repo_id=repo_id, revision=revision, files_metadata=True)38    remote_paths = {sibling.rfilename for sibling in info.siblings}39    resolved_revision = str(info.sha or revision or "main")40    if SNAPSHOTS_LATEST_PATH in remote_paths:41        return _materialize_hf_snapshot_repo_snapshot(42            repo_id=repo_id,43            local_dir=local_dir,44            revision=resolved_revision,45            requested_revision=revision,46            hf_sha=info.sha,47            remote_paths=remote_paths,48        )49    if {"issues.parquet", "pull_requests.parquet"} <= remote_paths:50        return _materialize_hf_root_snapshot(51            repo_id=repo_id,52            local_dir=local_dir,53            revision=resolved_revision,54            requested_revision=revision,55            hf_sha=info.sha,56            remote_paths=remote_paths,57        )58    return _materialize_hf_dataset_viewer_snapshot(59        repo_id=repo_id,60        local_dir=local_dir,61        revision=resolved_revision,62        requested_revision=revision,63        hf_sha=info.sha,64    )65 66 67def _materialize_hf_snapshot_repo_snapshot(68    *,69    repo_id: str,70    local_dir: Path,71    revision: str,72    requested_revision: str | None,73    hf_sha: str | None,74    remote_paths: set[str],75) -> Path:76    local_dir.mkdir(parents=True, exist_ok=True)77    latest_download = Path(78        hf_hub_download(79            repo_id=repo_id,80            repo_type="dataset",81            filename=SNAPSHOTS_LATEST_PATH,82            revision=revision,83        )84    )85    latest_payload = json.loads(latest_download.read_text(encoding="utf-8"))86    downloaded_files: set[str] = set()87    _copy_downloaded_file(88        latest_download, repo_relative_path_to_local(local_dir, SNAPSHOTS_LATEST_PATH)89    )90    downloaded_files.add(SNAPSHOTS_LATEST_PATH)91 92    for filename in (93        *RAW_TABLE_FILENAMES,94        ROOT_MANIFEST_FILENAME,95        PR_SCOPE_CLUSTERS_FILENAME,96        *CONTRIBUTOR_ARTIFACT_FILENAMES,97        *LEGACY_ANALYSIS_FILENAMES,98    ):99        downloaded = _download_first_available_hf_file(100            repo_id=repo_id,101            revision=revision,102            filenames=_hf_latest_snapshot_candidates(latest_payload, filename),103        )104        if downloaded is None:105            continue106        _copy_downloaded_file(downloaded, local_dir / filename)107        downloaded_files.add(filename)108 109    if STATE_WATERMARK_PATH in remote_paths:110        _download_repo_file(111            repo_id=repo_id,112            revision=revision,113            local_dir=local_dir,114            repo_path=STATE_WATERMARK_PATH,115            downloaded_files=downloaded_files,116        )117 118    _download_analysis_state_files(119        repo_id=repo_id,120        revision=revision,121        local_dir=local_dir,122        remote_paths=remote_paths,123        downloaded_files=downloaded_files,124    )125 126    _download_published_analysis_files(127        repo_id=repo_id,128        revision=revision,129        local_dir=local_dir,130        remote_paths=remote_paths,131        downloaded_files=downloaded_files,132    )133 134    _download_repo_file(135        repo_id=repo_id,136        revision=revision,137        local_dir=local_dir,138        repo_path=README_FILENAME,139        downloaded_files=downloaded_files,140        required=False,141    )142 143    manifest = (144        read_json(local_dir / ROOT_MANIFEST_FILENAME)145        if (local_dir / ROOT_MANIFEST_FILENAME).exists()146        else {}147    )148    manifest.setdefault("repo", _infer_repo_from_materialized_snapshot(local_dir))149    manifest.setdefault(150        "snapshot_id",151        str(latest_payload.get("latest_snapshot_id") or hf_sha or local_dir.name),152    )153    manifest.update(154        {155            "source_type": "hf_snapshot_repo",156            "hf_repo_id": repo_id,157            "hf_revision": requested_revision,158            "hf_resolved_revision": revision,159            "hf_sha": hf_sha,160            "materialized_at": _iso_now(),161            "downloaded_files": sorted(downloaded_files),162            "hf_latest_pointer": latest_payload,163        }164    )165    write_text(json.dumps(manifest, indent=2) + "\n", local_dir / ROOT_MANIFEST_FILENAME)166    return local_dir167 168 169def _materialize_hf_root_snapshot(170    *,171    repo_id: str,172    local_dir: Path,173    revision: str,174    requested_revision: str | None,175    hf_sha: str | None,176    remote_paths: set[str],177) -> Path:178    local_dir.mkdir(parents=True, exist_ok=True)179    downloaded_files: set[str] = set()180    for repo_path in (181        *RAW_TABLE_FILENAMES,182        ROOT_MANIFEST_FILENAME,183        PR_SCOPE_CLUSTERS_FILENAME,184        *CONTRIBUTOR_ARTIFACT_FILENAMES,185        *LEGACY_ANALYSIS_FILENAMES,186        SNAPSHOTS_LATEST_PATH,187        STATE_WATERMARK_PATH,188        README_FILENAME,189    ):190        if repo_path not in remote_paths:191            continue192        _download_repo_file(193            repo_id=repo_id,194            revision=revision,195            local_dir=local_dir,196            repo_path=repo_path,197            downloaded_files=downloaded_files,198        )199 200    _download_analysis_state_files(201        repo_id=repo_id,202        revision=revision,203        local_dir=local_dir,204        remote_paths=remote_paths,205        downloaded_files=downloaded_files,206    )207 208    _download_published_analysis_files(209        repo_id=repo_id,210        revision=revision,211        local_dir=local_dir,212        remote_paths=remote_paths,213        downloaded_files=downloaded_files,214    )215 216    manifest = (217        read_json(local_dir / ROOT_MANIFEST_FILENAME)218        if (local_dir / ROOT_MANIFEST_FILENAME).exists()219        else {}220    )221    manifest.setdefault("repo", _infer_repo_from_materialized_snapshot(local_dir))222    manifest.setdefault("snapshot_id", hf_sha or local_dir.name)223    manifest.update(224        {225            "source_type": "hf_root_snapshot",226            "hf_repo_id": repo_id,227            "hf_revision": requested_revision,228            "hf_resolved_revision": revision,229            "hf_sha": hf_sha,230            "materialized_at": _iso_now(),231            "downloaded_files": sorted(downloaded_files),232        }233    )234    write_text(json.dumps(manifest, indent=2) + "\n", local_dir / ROOT_MANIFEST_FILENAME)235    return local_dir236 237 238def _materialize_hf_dataset_viewer_snapshot(239    *,240    repo_id: str,241    local_dir: Path,242    revision: str,243    requested_revision: str | None,244    hf_sha: str | None,245) -> Path:246    local_dir.mkdir(parents=True, exist_ok=True)247    downloaded_files: set[str] = set()248    for index, url in enumerate(_hf_dataset_parquet_urls(repo_id, revision)):249        temporary_path = local_dir / f"tmp-{index:04d}.parquet"250        _download_url_to_path(url, temporary_path)251        table_name = _parquet_table_name(temporary_path)252        temporary_path.replace(local_dir / table_name)253        downloaded_files.add(table_name)254 255    readme_path = hf_hub_download(256        repo_id=repo_id,257        repo_type="dataset",258        filename=README_FILENAME,259        revision=revision,260    )261    shutil.copy2(readme_path, local_dir / README_FILENAME)262    downloaded_files.add(README_FILENAME)263    manifest = {264        "repo": _infer_repo_from_materialized_snapshot(local_dir),265        "snapshot_id": hf_sha or local_dir.name,266        "source_type": "hf_dataset_viewer",267        "hf_repo_id": repo_id,268        "hf_revision": requested_revision,269        "hf_resolved_revision": revision,270        "hf_sha": hf_sha,271        "materialized_at": _iso_now(),272        "downloaded_files": sorted(downloaded_files),273    }274    write_text(json.dumps(manifest, indent=2) + "\n", local_dir / ROOT_MANIFEST_FILENAME)275    return local_dir276 277 278def _download_published_analysis_files(279    *,280    repo_id: str,281    revision: str,282    local_dir: Path,283    remote_paths: set[str],284    downloaded_files: set[str],285) -> None:286    if CURRENT_ANALYSIS_MANIFEST_PATH in remote_paths:287        manifest_path = _download_repo_file(288            repo_id=repo_id,289            revision=revision,290            local_dir=local_dir,291            repo_path=CURRENT_ANALYSIS_MANIFEST_PATH,292            downloaded_files=downloaded_files,293        )294        current_manifest = load_current_analysis_manifest(manifest_path)295        for repo_path in _manifest_artifact_paths(current_manifest, include_archived=True):296            if repo_path not in remote_paths:297                continue298            _download_repo_file(299                repo_id=repo_id,300                revision=revision,301                local_dir=local_dir,302                repo_path=repo_path,303                downloaded_files=downloaded_files,304            )305 306    for repo_path in sorted(307        path for path in remote_paths if _is_archived_analysis_manifest_path(path)308    ):309        manifest_path = _download_repo_file(310            repo_id=repo_id,311            revision=revision,312            local_dir=local_dir,313            repo_path=repo_path,314            downloaded_files=downloaded_files,315        )316        archived_manifest = load_archived_analysis_run_manifest(manifest_path)317        for artifact_path in _manifest_artifact_paths(archived_manifest, include_archived=False):318            if artifact_path not in remote_paths:319                continue320            _download_repo_file(321                repo_id=repo_id,322                revision=revision,323                local_dir=local_dir,324                repo_path=artifact_path,325                downloaded_files=downloaded_files,326            )327 328 329def _download_analysis_state_files(330    *,331    repo_id: str,332    revision: str,333    local_dir: Path,334    remote_paths: set[str],335    downloaded_files: set[str],336) -> None:337    for repo_path in sorted(338        path for path in remote_paths if PurePosixPath(path).parts[:1] == ("analysis-state",)339    ):340        _download_repo_file(341            repo_id=repo_id,342            revision=revision,343            local_dir=local_dir,344            repo_path=repo_path,345            downloaded_files=downloaded_files,346        )347 348 349def _manifest_artifact_paths(350    payload: dict[str, Any],351    *,352    include_archived: bool,353) -> list[str]:354    paths = [355        str(value) for value in (payload.get("artifacts") or {}).values() if isinstance(value, str)356    ]357    if include_archived:358        paths.extend(359            str(value)360            for value in (payload.get("archived_artifacts") or {}).values()361            if isinstance(value, str)362        )363    deduped: list[str] = []364    seen: set[str] = set()365    for repo_path in paths:366        normalized = repo_path.lstrip("./")367        if not normalized or normalized in seen:368            continue369        seen.add(normalized)370        deduped.append(normalized)371    return deduped372 373 374def _is_archived_analysis_manifest_path(repo_path: str) -> bool:375    parts = PurePosixPath(repo_path).parts376    return (377        len(parts) == 5378        and parts[0] == "snapshots"379        and parts[2] == "analysis-runs"380        and parts[4] == ROOT_MANIFEST_FILENAME381    )382 383 384def _download_repo_file(385    *,386    repo_id: str,387    revision: str,388    local_dir: Path,389    repo_path: str,390    downloaded_files: set[str],391    required: bool = True,392) -> Path:393    try:394        downloaded = Path(395            hf_hub_download(396                repo_id=repo_id,397                repo_type="dataset",398                filename=repo_path,399                revision=revision,400            )401        )402    except Exception:403        if required:404            raise405        return local_dir / repo_path406    destination = repo_relative_path_to_local(local_dir, repo_path)407    _copy_downloaded_file(downloaded, destination)408    downloaded_files.add(repo_path)409    return destination410 411 412def _copy_downloaded_file(downloaded_path: Path, destination: Path) -> None:413    destination.parent.mkdir(parents=True, exist_ok=True)414    shutil.copy2(downloaded_path, destination)415 416 417def _hf_dataset_info(repo_id: str, revision: str | None, *, files_metadata: bool) -> Any:418    api = HfApi()419    try:420        return api.dataset_info(repo_id=repo_id, revision=revision, files_metadata=files_metadata)421    except TypeError:422        return api.dataset_info(repo_id=repo_id, revision=revision)423 424 425def _hf_dataset_parquet_urls(repo_id: str, revision: str | None = None) -> list[str]:426    query = urllib.parse.urlencode({"revision": revision}) if revision else ""427    api_url = (428        f"https://huggingface.co/api/datasets/{urllib.parse.quote(repo_id, safe='')}/parquet"429        f"{f'?{query}' if query else ''}"430    )431    with urlopen_with_retry(api_url, timeout=120, label=api_url) as response:432        payload = json.loads(response.read().decode("utf-8"))433    urls = payload.get("default", {}).get("train", [])434    if not isinstance(urls, list) or not urls:435        raise FileNotFoundError(436            f"No parquet export URLs found for HF dataset {repo_id} at {api_url}"437        )438    return [str(url) for url in urls]439 440 441def _download_first_available_hf_file(442    *,443    repo_id: str,444    revision: str,445    filenames: list[str],446) -> Path | None:447    for filename in filenames:448        try:449            downloaded = Path(450                hf_hub_download(451                    repo_id=repo_id,452                    repo_type="dataset",453                    filename=filename,454                    revision=revision,455                )456            )457        except Exception:458            continue459        if downloaded.exists():460            return downloaded461    return None462 463 464def _hf_latest_snapshot_candidates(latest_payload: dict[str, Any], filename: str) -> list[str]:465    candidates: list[str] = []466    manifest_path = str(latest_payload.get("manifest_path") or "").strip("/")467    snapshot_dir = str(latest_payload.get("snapshot_dir") or "").strip("/")468    latest_snapshot_id = str(latest_payload.get("latest_snapshot_id") or "").strip()469    archived_manifest_path = str(latest_payload.get("archived_manifest_path") or "").strip("/")470 471    if filename == ROOT_MANIFEST_FILENAME and manifest_path:472        candidates.append(manifest_path)473    if snapshot_dir and snapshot_dir not in {".", "/"}:474        candidates.append(f"{snapshot_dir}/{filename}")475    if filename == ROOT_MANIFEST_FILENAME and archived_manifest_path:476        candidates.append(archived_manifest_path)477    if manifest_path and "/" in manifest_path:478        manifest_dir = manifest_path.rsplit("/", 1)[0]479        candidates.append(f"{manifest_dir}/{filename}")480    if latest_snapshot_id:481        candidates.append(str(PurePosixPath("snapshots") / latest_snapshot_id / filename))482    candidates.append(filename)483 484    seen: set[str] = set()485    deduped: list[str] = []486    for candidate in candidates:487        normalized = candidate.lstrip("./")488        if not normalized or normalized in seen:489            continue490        seen.add(normalized)491        deduped.append(normalized)492    return deduped493 494 495def _download_url_to_path(url: str, destination: Path) -> None:496    destination.parent.mkdir(parents=True, exist_ok=True)497    urllib.request.urlretrieve(url, destination)498 499 500def _parquet_table_name(path: Path) -> str:501    import pyarrow.parquet as pq502 503    columns = set(pq.read_table(path).column_names)504    if {"parent_kind", "issue_api_url", "body"} <= columns:505        return "comments.parquet"506    if {"event", "source_issue_number", "source_issue_url"} <= columns:507        return "events.parquet"508    if {"milestone_title", "comments_count"} <= columns and "merged_at" not in columns:509        return "issues.parquet"510    if {"link_type", "link_origin", "target_number"} <= columns:511        return "links.parquet"512    if {"pull_request_number", "filename", "blob_url", "patch"} <= columns:513        return "pr_files.parquet"514    if {"pull_request_number", "diff", "html_url", "api_url"} <= columns:515        return "pr_diffs.parquet"516    if {"merged_at", "head_ref", "base_ref"} <= columns:517        return "pull_requests.parquet"518    if {"review_id", "pull_request_api_url", "path"} <= columns:519        return "review_comments.parquet"520    if {"pull_request_number", "submitted_at"} <= columns and "review_id" not in columns:521        return "reviews.parquet"522    raise ValueError(f"Unrecognized HF parquet schema for {path.name}: {sorted(columns)}")523 524 525def _infer_repo_from_materialized_snapshot(local_dir: Path) -> str:526    import pyarrow.parquet as pq527 528    for table_filename in RAW_TABLE_FILENAMES:529        path = local_dir / table_filename530        if not path.exists():531            continue532        rows = pq.read_table(path).slice(0, 1).to_pylist()533        if rows and rows[0].get("repo"):534            return str(rows[0]["repo"])535    raise FileNotFoundError(f"Could not infer repo from materialized snapshot in {local_dir}")536 537 538def _iso_now() -> str:539    return datetime.now(tz=UTC).replace(microsecond=0).isoformat().replace("+00:00", "Z")540