Safetensors
GeoText1652_model / server.py
mhassanch's picture
added hub storage and patch embeddings
d35edd7
Raw History Blame Contribute Delete
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"
)
@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