"""HTTP server for a GeoText-1652 custom Inference Endpoint container.""" from __future__ import annotations import os from contextlib import asynccontextmanager from pathlib import Path from typing import Any, Dict, List, Optional, Tuple import numpy as np import torch from fastapi import FastAPI, HTTPException from hub_storage import persist_patch_embeddings from infer import fetch_rgb_geotiff, load_model, model_outputs, tile_image from pydantic import BaseModel, Field class InferenceRequest(BaseModel): image_url: str = Field(..., description="Public RGB GeoTIFF/COG URL") queries: List[str] = Field(..., min_items=1, description="Text queries to rank") include_embeddings: bool = Field( False, description="Return 256-dimensional image and text vectors" ) include_bboxes: bool = Field(False, description="Return text-conditioned normalized boxes") include_patch_embeddings: bool = Field( False, description="Return normalized 256-dimensional spatial patch embeddings", ) store_patch_embeddings: bool = Field( False, description="Persist patch embeddings as Parquet in the configured Hub bucket" ) output_prefix: Optional[str] = Field( None, description="Optional Hub bucket key prefix for persisted patch embeddings" ) tile_overlap: int = Field(64, ge=0, description="Source-pixel overlap between tiles") tile_size: Optional[int] = Field( None, ge=64, description="Optional source-pixel tile size for tiled inference" ) return_tile_results: bool = Field( False, description="Include per-tile scores and regions in addition to stitched results" ) @asynccontextmanager async def lifespan(app: FastAPI): device = torch.device("cuda" if torch.cuda.is_available() else "cpu") model_dir = Path(os.environ.get("GEOTEXT_MODEL_DIR", "/repository")) source_dir = Path(os.environ.get("GEOTEXT_SOURCE_DIR", "/app/GeoText-1652")) model, tokenizer, config = load_model(source_dir, model_dir, device) app.state.runtime = { "config": config, "device": device, "model": model, "tokenizer": tokenizer, } yield app = FastAPI(title="GeoText-1652", lifespan=lifespan) def _pixel_box(box: List[float], tile: Dict[str, Any]) -> List[float]: """Convert a tile-local normalized cx/cy/w/h box into source pixel xyxy.""" cx, cy, width, height = box x1 = max(tile["column"], tile["column"] + (cx - width / 2) * tile["width"]) y1 = max(tile["row"], tile["row"] + (cy - height / 2) * tile["height"]) x2 = min(tile["column"] + tile["width"], tile["column"] + (cx + width / 2) * tile["width"]) y2 = min(tile["row"] + tile["height"], tile["row"] + (cy + height / 2) * tile["height"]) return [float(x1), float(y1), float(x2), float(y2)] def _iou(left: List[float], right: List[float]) -> float: x1 = max(left[0], right[0]) y1 = max(left[1], right[1]) x2 = min(left[2], right[2]) y2 = min(left[3], right[3]) intersection = max(0.0, x2 - x1) * max(0.0, y2 - y1) if not intersection: return 0.0 left_area = max(0.0, left[2] - left[0]) * max(0.0, left[3] - left[1]) right_area = max(0.0, right[2] - right[0]) * max(0.0, right[3] - right[1]) return intersection / max(left_area + right_area - intersection, 1e-8) def _stitch_boxes( tile_outputs: List[Tuple[Dict[str, Any], Dict[str, Any]]], image_width: int, image_height: int ) -> List[Dict[str, Any]]: candidates = [] for tile, output in tile_outputs: score_by_query = {item["query"]: item["similarity"] for item in output["ranked_queries"]} for region in output.get("text_conditioned_boxes", []): pixel_box = _pixel_box(region["box"], tile) candidates.append( { "box": pixel_box, "query": region["query"], "score": score_by_query[region["query"]], "tile_index": tile["index"], } ) stitched = [] for query in sorted({candidate["query"] for candidate in candidates}): pending = sorted( (candidate for candidate in candidates if candidate["query"] == query), key=lambda candidate: candidate["score"], reverse=True, ) while pending: selected = pending.pop(0) stitched.append( { "query": query, "score": selected["score"], "source_pixel_xyxy": selected["box"], "source_normalized_xyxy": [ selected["box"][0] / image_width, selected["box"][1] / image_height, selected["box"][2] / image_width, selected["box"][3] / image_height, ], "tile_index": selected["tile_index"], } ) pending = [ candidate for candidate in pending if _iou(selected["box"], candidate["box"]) < 0.5 ] return stitched def _tiled_outputs( runtime: Dict[str, Any], image, request: InferenceRequest ) -> Dict[str, Any]: tiles = tile_image(image, request.tile_size, request.tile_overlap) tile_outputs = [] for tile in tiles: output = model_outputs( runtime["model"], runtime["tokenizer"], runtime["config"], tile["image"], request.queries, runtime["device"], request.include_embeddings, request.include_bboxes, request.include_patch_embeddings or request.store_patch_embeddings, ) tile_outputs.append((tile, output)) best_by_query = {} for tile, output in tile_outputs: for item in output["ranked_queries"]: candidate = dict(item, tile_index=tile["index"]) if ( item["query"] not in best_by_query or item["similarity"] > best_by_query[item["query"]]["similarity"] ): best_by_query[item["query"]] = candidate result = { "ranked_queries": sorted( best_by_query.values(), key=lambda item: item["similarity"], reverse=True ), "tiling": { "tile_size": request.tile_size, "tile_overlap": request.tile_overlap, "tile_count": len(tiles), "similarity_stitching": "maximum tile similarity per query", }, } if request.include_embeddings: weights = np.asarray( [tile["width"] * tile["height"] for tile, _ in tile_outputs], dtype=np.float32 ) vectors = np.asarray( [output["image_embedding"] for _, output in tile_outputs], dtype=np.float32 ) image_embedding = np.average(vectors, axis=0, weights=weights) image_embedding /= max(float(np.linalg.norm(image_embedding)), 1e-8) result["image_embedding"] = image_embedding.tolist() result["text_embeddings"] = tile_outputs[0][1]["text_embeddings"] if request.include_bboxes: result["stitched_text_conditioned_boxes"] = _stitch_boxes( tile_outputs, image.width, image.height ) if request.include_patch_embeddings or request.store_patch_embeddings: patches = [] for tile, output in tile_outputs: for patch in output["patch_embeddings"]["patches"]: x1, y1, x2, y2 = patch["source_pixel_xyxy"] patches.append( { **patch, "tile_index": tile["index"], "source_pixel_xyxy": [ x1 + tile["column"], y1 + tile["row"], x2 + tile["column"], y2 + tile["row"], ], } ) result["patch_embeddings"] = { "embedding_dimension": tile_outputs[0][1]["patch_embeddings"]["embedding_dimension"], "grid": "per_tile", "patches": patches, } if request.return_tile_results: result["tile_results"] = [ { "tile_index": tile["index"], "source_pixel_window": [ tile["column"], tile["row"], tile["width"], tile["height"], ], **output, } for tile, output in tile_outputs ] return result @app.get("/health") def health() -> Dict[str, str]: return {"status": "ok"} @app.post("/") @app.post("/infer") def infer(request: InferenceRequest) -> Dict[str, Any]: try: image, image_metadata = fetch_rgb_geotiff(request.image_url) runtime = app.state.runtime result = { "image_url": request.image_url, "image_metadata": image_metadata, "device": str(runtime["device"]), } include_patches = request.include_patch_embeddings or request.store_patch_embeddings if request.tile_size is None: result.update( model_outputs( runtime["model"], runtime["tokenizer"], runtime["config"], image, request.queries, runtime["device"], include_embeddings=request.include_embeddings, include_bboxes=request.include_bboxes, include_patch_embeddings=include_patches, ) ) else: result.update(_tiled_outputs(runtime, image, request)) if request.store_patch_embeddings: result["patch_embeddings_storage"] = persist_patch_embeddings( image_url=request.image_url, patches=result["patch_embeddings"]["patches"], output_prefix=request.output_prefix, ) if not request.include_patch_embeddings: result.pop("patch_embeddings") return result except Exception as error: raise HTTPException(status_code=422, detail=str(error)) from error