Download server.py from geobase/GeoText1652_model: direct link, hf CLI and curl.
- Browser
- Download file 10.3 kB
-
https://huggingface.co/geobase/GeoText1652_model/resolve/main/server.py
- Command line
-
hf download hf://geobase/GeoText1652_model/server.py
-
curl -L -o server.py https://huggingface.co/geobase/GeoText1652_model/resolve/main/server.py
10.3 kB
| """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" | |
| ) | |
| 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 | |
| def health() -> Dict[str, str]: | |
| return {"status": "ok"} | |
| 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 | |