"""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 ), }