"""Persist GeoText patch embeddings in a Hugging Face dataset repository.""" from __future__ import annotations import os import tempfile import uuid from datetime import datetime, timezone from typing import Any, Dict, List import pyarrow as pa import pyarrow.parquet as pq from huggingface_hub import HfApi def _repository_config() -> Dict[str, str]: token = os.environ.get("HF_TOKEN") or os.environ.get("HUGGING_FACE_HUB_TOKEN") repository = os.environ.get("HF_BUCKET") if not token or not repository or "/" not in repository: raise ValueError( "store_patch_embeddings=true requires HF_TOKEN (or " "HUGGING_FACE_HUB_TOKEN) and HF_BUCKET=namespace/repository-name" ) return {"token": token, "repository": repository} def persist_patch_embeddings( *, image_url: str, patches: List[Dict[str, Any]], output_prefix: str | None ) -> Dict[str, Any]: """Write patch vectors as Parquet and upload them through the stable Hub API.""" config = _repository_config() schema = pa.schema( [ pa.field("id", pa.string()), pa.field("image_url", pa.string()), pa.field("tile_index", pa.int32()), pa.field("patch_index", pa.int32()), pa.field("row", pa.int32()), pa.field("column", pa.int32()), pa.field("source_pixel_xyxy", pa.list_(pa.float32())), pa.field("embedding", pa.list_(pa.float32())), ] ) rows = [ { "id": f"tile_{patch.get('tile_index', 0)}_patch_{patch['patch_index']:04d}", "image_url": image_url, "tile_index": patch.get("tile_index", 0), "patch_index": patch["patch_index"], "row": patch["row"], "column": patch["column"], "source_pixel_xyxy": patch["source_pixel_xyxy"], "embedding": patch["embedding"], } for patch in patches ] table = pa.Table.from_pylist(rows, schema=schema) date = datetime.now(timezone.utc).strftime("%Y%m%d") prefix = output_prefix or f"geotext-patch-embeddings/{date}/{uuid.uuid4().hex}/" key = f"{prefix.rstrip('/')}/patch_embeddings.parquet" with tempfile.NamedTemporaryFile(suffix=".parquet") as artifact: pq.write_table(table, artifact.name, compression="zstd") api = HfApi(token=config["token"]) api.create_repo( repo_id=config["repository"], repo_type="dataset", private=True, exist_ok=True ) api.upload_file( path_or_fileobj=artifact.name, path_in_repo=key, repo_id=config["repository"], repo_type="dataset", commit_message="Store GeoText patch embeddings", ) return { "provider": "huggingface_hub", "repository": config["repository"], "key": key, "format": "parquet", "row_count": len(rows), }