File size: 1,930 Bytes
4d3c316 6382bf0 4d3c316 6382bf0 4d3c316 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 | """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
),
}
|