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