Safetensors
GeoText1652_model / test_endpoint.py
mhassanch's picture
Support JPEG and PNG image inputs
ac44754
Raw History Blame Contribute Delete
4.33 kB
#!/usr/bin/env python3
"""Smoke-test a deployed GeoText-1652 Inference Endpoint.
Usage:
export HF_TOKEN=hf_...
python test_endpoint.py
Optional:
GEOTEXT_ENDPOINT_URL=https://...endpoints.huggingface.cloud \
python test_endpoint.py --patches --store-patches
"""
from __future__ import annotations
import argparse
import os
import sys
from typing import Any, Dict, Optional
import requests
DEFAULT_ENDPOINT = "https://6ac5162dc8b01bf6eda56cb3.endpoints.huggingface.cloud"
DEFAULT_IMAGE = (
"https://huggingface.co/datasets/geobase/geoai-cogs/resolve/main/"
"geoembeddings-demo/building-detection_sm.tif"
)
def parse_args() -> argparse.Namespace:
parser = argparse.ArgumentParser(description=__doc__)
parser.add_argument(
"--endpoint-url",
default=os.environ.get("GEOTEXT_ENDPOINT_URL", DEFAULT_ENDPOINT),
help="Endpoint base URL; defaults to GEOTEXT_ENDPOINT_URL or the deployed endpoint.",
)
parser.add_argument("--image-url", default=DEFAULT_IMAGE, help="Public RGB GeoTIFF/COG URL.")
parser.add_argument(
"--patches", action="store_true", help="Also test tiled patch-embedding output."
)
parser.add_argument(
"--store-patches",
action="store_true",
help="Also test Parquet persistence; endpoint must have HF_TOKEN and HF_BUCKET configured.",
)
return parser.parse_args()
def call(
session: requests.Session,
method: str,
url: str,
payload: Optional[Dict[str, Any]] = None,
) -> Dict[str, Any]:
response = session.request(method, url, json=payload, timeout=300)
try:
response.raise_for_status()
except requests.HTTPError as error:
raise RuntimeError(f"{method} {url} failed: {response.text}") from error
return response.json()
def check(name: str, result: Dict[str, Any], required: set[str]) -> None:
missing = required.difference(result)
if missing:
raise RuntimeError(f"{name} response is missing: {sorted(missing)}\n{result}")
print(f"PASS {name}")
def main() -> None:
args = parse_args()
token = os.environ.get("HF_TOKEN")
if not token:
sys.exit("Set HF_TOKEN to a Hugging Face token with access to the private endpoint.")
base_url = args.endpoint_url.rstrip("/")
session = requests.Session()
session.headers.update(
{"Authorization": f"Bearer {token}", "Content-Type": "application/json"}
)
health = call(session, "GET", f"{base_url}/health")
check("health", health, {"status"})
text = call(
session,
"POST",
f"{base_url}/embed/text",
{"queries": ["a parking lot", "a building"]},
)
check("text embeddings", text, {"embedding_dimension", "text_embeddings"})
assert len(text["text_embeddings"]) == 2
assert len(text["text_embeddings"][0]["embedding"]) == 256
image = call(
session,
"POST",
f"{base_url}/embed/image",
{"image_url": args.image_url, "include_global_embedding": True},
)
check("global image embedding", image, {"image_embedding", "image_metadata"})
assert len(image["image_embedding"]) == 256
ranking = call(
session,
"POST",
f"{base_url}/infer",
{"image_url": args.image_url, "queries": ["a parking lot", "a building"]},
)
check("text-guided ranking", ranking, {"ranked_queries", "image_metadata"})
assert len(ranking["ranked_queries"]) == 2
if args.patches or args.store_patches:
patch_request: Dict[str, Any] = {
"image_url": args.image_url,
"tile_size": 384,
"tile_overlap": 64,
"include_global_embedding": True,
"include_patch_embeddings": args.patches,
"store_patch_embeddings": args.store_patches,
}
patches = call(session, "POST", f"{base_url}/embed/image", patch_request)
check("tiled patch embeddings", patches, {"tiling"})
if args.patches:
check("patch-vector response", patches, {"patch_embeddings"})
assert patches["patch_embeddings"]["patches"]
if args.store_patches:
check("patch-vector storage", patches, {"patch_embeddings_storage"})
print("All requested GeoText endpoint checks passed.")
if __name__ == "__main__":
main()