Safetensors
GeoText1652_model / server.py
mhassanch's picture
Support JPEG and PNG image inputs
ac44754
Raw History Blame Contribute Delete
16.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,
image_embedding_outputs,
load_model,
model_outputs,
text_embedding_outputs,
tile_image,
)
from pydantic import BaseModel, Field
class InferenceRequest(BaseModel):
image_url: str = Field(..., description="Public RGB GeoTIFF/COG, JPEG, or PNG 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"
)
class TextEmbeddingRequest(BaseModel):
queries: List[str] = Field(..., min_items=1, description="Text strings to embed")
class ImageEmbeddingRequest(BaseModel):
image_url: str = Field(..., description="Public RGB GeoTIFF/COG, JPEG, or PNG URL")
include_global_embedding: bool = Field(
True, description="Return the normalized 256-dimensional image vector"
)
include_patch_embeddings: bool = Field(
False, description="Return normalized 256-dimensional spatial patch vectors"
)
store_patch_embeddings: bool = Field(
False, description="Persist patch embeddings as Parquet in HF_BUCKET"
)
output_prefix: Optional[str] = Field(
None, description="Optional Hugging Face dataset path prefix"
)
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"
)
return_tile_results: bool = Field(
False, description="Include per-tile embeddings in addition to stitched output"
)
@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
def _image_outputs(
runtime: Dict[str, Any], image, request: ImageEmbeddingRequest
) -> Dict[str, Any]:
"""Run vision-only embedding inference, optionally over overlapping tiles."""
include_patches = request.include_patch_embeddings or request.store_patch_embeddings
if request.tile_size is None:
return image_embedding_outputs(
runtime["model"],
runtime["config"],
image,
runtime["device"],
include_global_embedding=request.include_global_embedding,
include_patch_embeddings=include_patches,
)
tiles = tile_image(image, request.tile_size, request.tile_overlap)
tile_outputs = []
for tile in tiles:
output = image_embedding_outputs(
runtime["model"],
runtime["config"],
tile["image"],
runtime["device"],
include_global_embedding=request.include_global_embedding,
include_patch_embeddings=include_patches,
)
tile_outputs.append((tile, output))
result = {
"embedding_dimension": tile_outputs[0][1]["embedding_dimension"],
"tiling": {
"tile_size": request.tile_size,
"tile_overlap": request.tile_overlap,
"tile_count": len(tiles),
"global_embedding_stitching": "area-weighted mean, then L2 normalization",
},
}
if request.include_global_embedding:
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
)
embedding = np.average(vectors, axis=0, weights=weights)
embedding /= max(float(np.linalg.norm(embedding)), 1e-8)
result["image_embedding"] = embedding.tolist()
if include_patches:
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": result["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]:
"""Backward-compatible combined image/text search endpoint."""
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
@app.post("/embed/text")
def embed_text(request: TextEmbeddingRequest) -> Dict[str, Any]:
"""Return query vectors without fetching imagery or running the vision encoder."""
try:
runtime = app.state.runtime
result = text_embedding_outputs(
runtime["model"],
runtime["tokenizer"],
runtime["config"],
request.queries,
runtime["device"],
)
return {"device": str(runtime["device"]), **result}
except Exception as error:
raise HTTPException(status_code=422, detail=str(error)) from error
@app.post("/embed/image")
def embed_image(request: ImageEmbeddingRequest) -> Dict[str, Any]:
"""Return image or patch vectors without tokenizing text or running search."""
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"]),
}
result.update(_image_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