Download hub_storage.py from geobase/GeoText1652_model: direct link, hf CLI and curl.
- Browser
- Download file 2.94 kB
-
https://huggingface.co/geobase/GeoText1652_model/resolve/main/hub_storage.py
- Command line
-
hf download hf://geobase/GeoText1652_model/hub_storage.py
-
curl -L -o hub_storage.py https://huggingface.co/geobase/GeoText1652_model/resolve/main/hub_storage.py
2.94 kB
| """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), | |
| } | |