Download handler.py from geobase/GeoText1652_model: direct link, hf CLI and curl.
- Browser
- Download file 1.93 kB
-
https://huggingface.co/geobase/GeoText1652_model/resolve/main/handler.py
- Command line
-
hf download hf://geobase/GeoText1652_model/handler.py
-
curl -L -o handler.py https://huggingface.co/geobase/GeoText1652_model/resolve/main/handler.py
1.93 kB
| """Hugging Face Inference Endpoints custom handler for GeoText-1652.""" | |
| from __future__ import annotations | |
| from pathlib import Path | |
| from typing import Any, Dict | |
| import torch | |
| from infer import fetch_rgb_geotiff, load_model, rank_queries | |
| class EndpointHandler: | |
| """Rank natural-language descriptions against one public RGB GeoTIFF/COG.""" | |
| def __init__(self, path: str = "") -> None: | |
| self.repository_dir = Path(path) if path else Path(__file__).resolve().parent | |
| self.device = torch.device("cuda" if torch.cuda.is_available() else "cpu") | |
| self.model, self.tokenizer, self.config = load_model( | |
| self.repository_dir / "GeoText-1652", self.repository_dir, self.device | |
| ) | |
| def __call__(self, data: Dict[str, Any]) -> Dict[str, Any]: | |
| inputs = data.get("inputs", data) | |
| if not isinstance(inputs, dict): | |
| raise ValueError("Request inputs must be a JSON object.") | |
| image_url = inputs.get("image_url") or inputs.get("image") | |
| if not isinstance(image_url, str) or not image_url: | |
| raise ValueError("inputs.image_url must be a public GeoTIFF/COG URL.") | |
| queries = inputs.get("queries", inputs.get("query")) | |
| if isinstance(queries, str): | |
| queries = [queries] | |
| if not isinstance(queries, list) or not queries or not all( | |
| isinstance(query, str) and query for query in queries | |
| ): | |
| raise ValueError("inputs.query or inputs.queries must contain one or more strings.") | |
| image, image_metadata = fetch_rgb_geotiff(image_url) | |
| return { | |
| "model": "truemanv5666/GeoText1652_model", | |
| "image_url": image_url, | |
| "image_metadata": image_metadata, | |
| "device": str(self.device), | |
| "ranked_queries": rank_queries( | |
| self.model, self.tokenizer, self.config, image, queries, self.device | |
| ), | |
| } | |