Safetensors
GeoText1652_model / hub_storage.py
mhassanch's picture
Refactor hub storage to use Hugging Face dataset repository
9da5d19
Raw History Blame Contribute Delete
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),
}