geospatial-api / inference.py
Rudraaaa76's picture
Update inference.py
339d60a verified
Raw History Blame Contribute Delete
24.4 kB
"""
Aerix Inference Module - SVAMITVA Feature Extraction
Performs building, road, and water body segmentation on orthophoto tiles using UNet++.
Fixed for ZeroGPU + large-TIFF stability:
1. Image loading uses rasterio windowed/decimated reads for GeoTIFFs, so a
multi-hundred-MB raster is never fully decoded into RAM just to be
downsampled. Falls back to PIL/OpenCV for ordinary PNG/JPG.
2. The pipeline is split into three explicit stages so app.py can keep the
CPU-only parts (file I/O, resizing, tiling, postprocessing) OUTSIDE the
@spaces.GPU-decorated function, and only the actual tensor->model->tensor
step inside it:
prepare_input() -> CPU
run_model() -> must run inside @spaces.GPU
postprocess() -> CPU
predict() / predict_large() still exist and just chain the three stages,
for local/non-ZeroGPU use (e.g. the CLI at the bottom of this file).
3. A cheap header-only size/dimension check (get_image_info) lets the caller
reject or route oversized uploads BEFORE attempting to open pixel data.
"""
import os
import sys
import cv2
import numpy as np
import torch
from pathlib import Path
import segmentation_models_pytorch as smp
from typing import Tuple, Dict, Optional, List
# Ensure UTF-8 output encoding for Windows consoles
if hasattr(sys.stdout, "reconfigure"):
try:
sys.stdout.reconfigure(encoding="utf-8")
except Exception:
pass
from PIL import Image
Image.MAX_IMAGE_PIXELS = None
try:
import rasterio
from rasterio.enums import Resampling as RioResampling
HAS_RASTERIO = True
except ImportError:
HAS_RASTERIO = False
_IMAGE_CACHE = {}
_CACHE_LIMIT = 8 # decimated images can still be several MB each; keep this small
# Hard ceiling on raw upload size. Tune to whatever your Space's hardware can
# actually hold in RAM alongside the model. Rejecting fast beats hanging.
# NOTE: this ceiling is safe well above the old 300MB specifically because
# GeoTIFFs go through rasterio's windowed read (STAGE 1 in prepare_input) and
# never get fully decoded into RAM regardless of file size. The remaining
# risk is non-TIFF uploads (PNG/JPG), which still fall back to a full PIL
# decode — a 1GB+ PNG could still spike RAM on a 16GB ZeroGPU container.
MAX_UPLOAD_BYTES = 1024 * 1024 * 1024 # 1 GB
def get_image_info(image_path: str) -> Dict:
"""
Header-only inspection — does NOT decode pixel data. Safe to call on a
multi-GB file. Used to size-gate uploads before touching pixels.
"""
size_bytes = os.path.getsize(image_path) if os.path.exists(image_path) else 0
width = height = None
if HAS_RASTERIO:
try:
with rasterio.open(image_path) as ds:
width, height = ds.width, ds.height
except Exception:
pass
if width is None:
try:
with Image.open(image_path) as pil_img:
width, height = pil_img.size
except Exception:
pass
return {"size_bytes": size_bytes, "width": width, "height": height}
def validate_upload(image_path: str, max_bytes: int = MAX_UPLOAD_BYTES) -> None:
"""Raise a clean, fast error instead of letting a huge file hang the pipeline."""
info = get_image_info(image_path)
if info["size_bytes"] > max_bytes:
raise ValueError(
f"File is {info['size_bytes'] / 1e6:.0f} MB, which exceeds the "
f"{max_bytes / 1e6:.0f} MB limit for this Space. Please downsample "
f"or crop the orthophoto before uploading."
)
def _load_via_rasterio(image_path: str, max_dim: int) -> np.ndarray:
"""
Decimated read: asks GDAL to decode directly at (approximately) the target
resolution using its own overviews/decimation, so the full-resolution
raster is never materialized in RAM. This is the key fix for large
GeoTIFFs — PIL/OpenCV have no equivalent of this.
"""
with rasterio.open(image_path) as ds:
w, h = ds.width, ds.height
scale = min(1.0, max_dim / float(max(w, h))) if max(w, h) > max_dim else 1.0
out_w, out_h = max(1, int(w * scale)), max(1, int(h * scale))
band_count = min(ds.count, 3)
data = ds.read(
indexes=list(range(1, band_count + 1)),
out_shape=(band_count, out_h, out_w),
resampling=RioResampling.bilinear,
) # (bands, H, W)
data = np.transpose(data, (1, 2, 0)) # HWC
if data.dtype != np.uint8:
# Common with GeoTIFFs (uint16 / float reflectance). Normalize per-band
# to 8-bit for the RGB model input.
data = data.astype(np.float32)
for b in range(data.shape[-1]):
band = data[..., b]
lo, hi = np.percentile(band, (1, 99))
if hi > lo:
band = np.clip((band - lo) / (hi - lo), 0, 1) * 255.0
else:
band = np.zeros_like(band)
data[..., b] = band
data = data.astype(np.uint8)
if data.shape[-1] == 1:
data = np.repeat(data, 3, axis=-1)
elif data.shape[-1] > 3:
data = data[..., :3]
return np.ascontiguousarray(data)
def _load_via_pil_or_cv2(image_path: str, max_dim: int) -> np.ndarray:
try:
with Image.open(image_path) as pil_img:
w, h = pil_img.size
if max(w, h) > max_dim:
scale = max_dim / float(max(w, h))
new_w, new_h = max(1, int(w * scale)), max(1, int(h * scale))
# draft() lets Pillow decode a smaller JPEG DCT block directly
# where supported; a no-op for formats that don't support it.
pil_img.draft("RGB", (new_w, new_h))
pil_img = pil_img.resize((new_w, new_h), Image.Resampling.BILINEAR)
if pil_img.mode != "RGB":
pil_img = pil_img.convert("RGB")
return np.array(pil_img)
except Exception:
image = cv2.imread(image_path)
if image is None:
raise ValueError(f"Cannot load image {image_path}. Ensure it is a valid PNG, JPG, or TIFF format.")
image = cv2.cvtColor(image, cv2.COLOR_BGR2RGB)
h, w = image.shape[:2]
if max(h, w) > max_dim:
scale = max_dim / float(max(h, w))
new_w, new_h = max(1, int(w * scale)), max(1, int(h * scale))
image = cv2.resize(image, (new_w, new_h))
return image
def load_robust_image(image_path: str, max_dim: int = 2048) -> np.ndarray:
"""
Robustly load any image format (PNG, JPG, TIFF, large GeoTIFF) as an RGB
numpy array, downsampled to max_dim on the longest side. For GeoTIFFs this
uses rasterio's decimated read (no full-res decode); everything else uses
PIL with draft-mode decoding, falling back to OpenCV.
"""
mtime = os.path.getmtime(image_path) if os.path.exists(image_path) else 0
cache_key = (str(image_path), max_dim, mtime)
if cache_key in _IMAGE_CACHE:
return _IMAGE_CACHE[cache_key].copy()
suffix = Path(image_path).suffix.lower()
result = None
if HAS_RASTERIO and suffix in (".tif", ".tiff"):
try:
result = _load_via_rasterio(image_path, max_dim)
except Exception as e:
print(f"rasterio read failed ({e}), falling back to PIL/OpenCV")
if result is None:
result = _load_via_pil_or_cv2(image_path, max_dim)
if len(_IMAGE_CACHE) >= _CACHE_LIMIT:
_IMAGE_CACHE.clear()
_IMAGE_CACHE[cache_key] = result
return result.copy()
class AerixSegmentationModel:
"""
Unified interface for SVAMITVA feature extraction using UNet++.
Supports buildings, roads, and water bodies segmentation.
Pipeline is split into three stages so the caller controls exactly what
runs under a ZeroGPU @spaces.GPU context:
prepare_input() — CPU only, safe to run with no GPU attached
run_model() — the only stage that needs CUDA
postprocess() — CPU only
"""
def __init__(self, models_dir: str = "models", device: str = None):
self.device = device or ("cuda" if torch.cuda.is_available() else "cpu")
self.models_dir = Path(models_dir)
self.models = {}
self.feature_names = ["buildings", "roads", "water_bodies"]
def _load_single_model(self, feature_name: str):
"""Load a single model on demand. Call this INSIDE the GPU context —
it's what does model.to('cuda') under ZeroGPU."""
if feature_name in self.models:
return self.models[feature_name]
model_path = self.models_dir / f"{feature_name}_unetpp_model.pth"
if not model_path.exists():
raise FileNotFoundError(f"Model weight file not found: {model_path}")
if len(self.models) > 0:
self.models.clear()
import gc
gc.collect()
if torch.cuda.is_available():
torch.cuda.empty_cache()
print(f"Loading {feature_name} model into RAM...")
try:
model = smp.UnetPlusPlus(
encoder_name="resnet34",
encoder_weights=None,
in_channels=3,
classes=1,
activation="sigmoid",
)
state_dict = torch.load(model_path, map_location=self.device, weights_only=True)
model.load_state_dict(state_dict)
model.to(self.device)
model.eval()
self.models[feature_name] = model
print(f"✓ Successfully loaded {feature_name} model")
return model
except Exception as e:
print(f"✗ Failed to load {feature_name}: {e}")
raise e
# ------------------------------------------------------------------ #
# STAGE 1 — CPU only. Safe to call before a GPU is attached.
# ------------------------------------------------------------------ #
def prepare_input(self, image_path: str, target_size: Tuple[int, int] = (512, 512),
max_dim: int = 2048) -> Dict:
"""Load + downsample the image and build the normalized model input array."""
validate_upload(image_path)
original_image = load_robust_image(image_path, max_dim=max_dim)
resized = cv2.resize(original_image, (target_size[1], target_size[0]))
input_array = resized.astype(np.float32) / 255.0 # HWC, CPU numpy — no torch/cuda touched
return {"original_image": original_image, "input_array": input_array}
# ------------------------------------------------------------------ #
# STAGE 2 — the ONLY stage that needs CUDA. Wrap the caller of this
# (or this itself) in @spaces.GPU on ZeroGPU Spaces.
# ------------------------------------------------------------------ #
def run_model(self, feature: str, input_array: np.ndarray) -> np.ndarray:
"""input_array: HWC float32 [0,1]. Returns raw sigmoid output, HxW numpy."""
model = self._load_single_model(feature)
tensor = torch.from_numpy(input_array).permute(2, 0, 1).unsqueeze(0).to(self.device)
with torch.no_grad():
output = model(tensor)
raw = output.squeeze().detach().cpu().numpy()
del tensor, output
return raw
# ------------------------------------------------------------------ #
# STAGE 3 — CPU only.
# ------------------------------------------------------------------ #
def postprocess(self, original_image: np.ndarray, raw_output: np.ndarray,
feature: str, threshold: float = 0.5) -> Dict:
binary_mask = (raw_output > threshold).astype(np.uint8) * 255
if original_image.shape[:2] != binary_mask.shape:
binary_mask = cv2.resize(binary_mask, (original_image.shape[1], original_image.shape[0]))
overlay = self._create_overlay(original_image, binary_mask)
detected_pixels = int(np.sum(binary_mask > 0))
total_pixels = int(binary_mask.size)
detection_ratio = detected_pixels / total_pixels if total_pixels > 0 else 0
rooftop_analysis = None
if feature == "buildings" and detected_pixels > 0:
rooftop_analysis = self.classify_rooftops(original_image, binary_mask)
return {
"original_image": original_image,
"mask": binary_mask,
"overlay": overlay,
"feature": feature,
"detected_ratio": detection_ratio,
"detected_pixels": detected_pixels,
"total_pixels": total_pixels,
"rooftop_analysis": rooftop_analysis,
}
# ------------------------------------------------------------------ #
# Convenience wrapper for local / non-ZeroGPU use (CLI, tests). On a
# ZeroGPU Space, DO NOT call this from inside app.py — call the three
# stages separately so only run_model() sits inside @spaces.GPU.
# ------------------------------------------------------------------ #
def predict(self, image_path: str, feature: str = "buildings", threshold: float = 0.5) -> Dict:
prepped = self.prepare_input(image_path)
raw = self.run_model(feature, prepped["input_array"])
return self.postprocess(prepped["original_image"], raw, feature, threshold)
@staticmethod
def classify_rooftops(image: np.ndarray, mask: np.ndarray) -> Dict:
"""
Rooftop classification (RCC, Tiled, Tin, Others) on detected building
footprints, plus solar/tax estimates. Pure CPU/OpenCV — no GPU needed.
"""
hsv = cv2.cvtColor(image, cv2.COLOR_RGB2HSV)
num_labels, labels, stats, centroids = cv2.connectedComponentsWithStats((mask > 0).astype(np.uint8))
rooftop_counts = {
"RCC / Concrete": 0,
"Tiled / Terracotta": 0,
"Tin / Metal Sheet": 0,
"Others": 0,
}
building_details = []
total_area_m2 = 0.0
rcc_area_m2 = 0.0
pixel_area_m2 = 0.25 # 50cm GSD -> 1 px = 0.25 m²
colored_classification_map = image.copy()
for i in range(1, num_labels):
area_px = stats[i, cv2.CC_STAT_AREA]
if area_px < 15:
continue
area_m2 = area_px * pixel_area_m2
total_area_m2 += area_m2
component_mask = (labels == i)
mean_hsv = cv2.mean(hsv, mask=component_mask.astype(np.uint8))
h, s, v = mean_hsv[0], mean_hsv[1], mean_hsv[2]
if (h <= 25 or h >= 160) and s > 35:
roof_type = "Tiled / Terracotta"
color = [220, 80, 50]
elif 90 <= h <= 135 and s > 30:
roof_type = "Tin / Metal Sheet"
color = [50, 120, 220]
elif s < 45 and v > 80:
roof_type = "RCC / Concrete"
color = [180, 180, 190]
rcc_area_m2 += area_m2
else:
roof_type = "Others"
color = [180, 140, 70]
rooftop_counts[roof_type] += 1
colored_classification_map[component_mask] = color
building_details.append({
"id": i,
"area_m2": round(area_m2, 2),
"roof_type": roof_type,
"centroid": (int(centroids[i][0]), int(centroids[i][1])),
})
usable_solar_m2 = rcc_area_m2 * 0.75
annual_solar_kwh = usable_solar_m2 * 150.0
tax_inr = (rcc_area_m2 * 25.0) + ((total_area_m2 - rcc_area_m2) * 15.0)
return {
"rooftop_counts": rooftop_counts,
"total_buildings": len(building_details),
"total_builtup_m2": round(total_area_m2, 2),
"rcc_area_m2": round(rcc_area_m2, 2),
"annual_solar_kwh": round(annual_solar_kwh, 1),
"annual_property_tax_inr": round(tax_inr, 2),
"classification_map": colored_classification_map,
"building_details": building_details,
}
def _create_overlay(self, image: np.ndarray, mask: np.ndarray, alpha: float = 0.55) -> np.ndarray:
overlay = image.copy()
binary = (mask > 127).astype(np.uint8)
if np.sum(binary) > 0:
color_fill = np.zeros_like(image)
color_fill[binary > 0] = [70, 180, 70]
overlay = cv2.addWeighted(image, 1.0 - alpha, color_fill, alpha, 0)
contours, _ = cv2.findContours(binary, cv2.RETR_EXTERNAL, cv2.CHAIN_APPROX_SIMPLE)
cv2.drawContours(overlay, contours, -1, (240, 255, 240), 1, cv2.LINE_AA)
cv2.drawContours(overlay, contours, -1, (40, 200, 40), 1, cv2.LINE_AA)
return overlay
# ------------------------------------------------------------------ #
# Large-orthomosaic path: sliding window, split the same way.
# ------------------------------------------------------------------ #
def prepare_tiles(self, image: np.ndarray, tile_size: int = 512,
overlap: int = 64) -> Dict:
"""CPU only: build the chip batch + blending weight map. No CUDA."""
h, w = image.shape[:2]
stride = tile_size - overlap
ramp = np.ones(tile_size, dtype=np.float32)
if overlap > 0:
fade = np.linspace(0, 1, overlap, dtype=np.float32)
ramp[:overlap] = fade
ramp[-overlap:] = fade[::-1]
weight_map = np.outer(ramp, ramp)
y_positions = list(range(0, max(1, h - tile_size + 1), stride))
if y_positions[-1] + tile_size < h:
y_positions.append(h - tile_size)
x_positions = list(range(0, max(1, w - tile_size + 1), stride))
if x_positions[-1] + tile_size < w:
x_positions.append(w - tile_size)
chips = []
coords = []
for y0 in y_positions:
for x0 in x_positions:
chip = image[y0:y0 + tile_size, x0:x0 + tile_size]
ch, cw = chip.shape[:2]
if ch < tile_size or cw < tile_size:
padded = np.zeros((tile_size, tile_size, 3), dtype=np.uint8)
padded[:ch, :cw] = chip
chip = padded
chips.append(chip.astype(np.float32) / 255.0)
coords.append((y0, x0, min(tile_size, h - y0), min(tile_size, w - x0)))
print(f" Sliding window: {len(y_positions)}×{len(x_positions)} = {len(chips)} tiles "
f"({tile_size}px, {overlap}px overlap)")
return {"chips": chips, "coords": coords, "weight_map": weight_map,
"image_shape": (h, w)}
def run_model_tiles(self, feature: str, chips: List[np.ndarray],
batch_size: Optional[int] = None) -> List[np.ndarray]:
"""The only CUDA-touching part of the large-image path."""
model = self._load_single_model(feature)
batch_size = batch_size or (4 if self.device == "cpu" else 16)
preds: List[np.ndarray] = []
with torch.no_grad():
for i in range(0, len(chips), batch_size):
batch = np.stack(chips[i:i + batch_size]) # (B, H, W, C)
tensor = torch.from_numpy(batch).permute(0, 3, 1, 2).to(self.device)
out = model(tensor).squeeze(1).detach().cpu().numpy()
if out.ndim == 2:
out = np.expand_dims(out, 0)
preds.extend(list(out))
del tensor, out
return preds
def stitch_tiles(self, preds: List[np.ndarray], coords, weight_map,
image_shape, threshold: float = 0.5) -> np.ndarray:
"""CPU only: blend overlapping tile predictions into one mask."""
h, w = image_shape
prob_accum = np.zeros((h, w), dtype=np.float64)
weight_accum = np.zeros((h, w), dtype=np.float64)
for pred, (y0, x0, ah, aw) in zip(preds, coords):
prob_accum[y0:y0 + ah, x0:x0 + aw] += pred[:ah, :aw] * weight_map[:ah, :aw]
weight_accum[y0:y0 + ah, x0:x0 + aw] += weight_map[:ah, :aw]
weight_accum[weight_accum == 0] = 1.0
prob_map = prob_accum / weight_accum
return (prob_map > threshold).astype(np.uint8) * 255
def predict_large(self, image_path: str, feature: str = "buildings",
threshold: float = 0.5, max_display: int = 2048,
tile_size: int = 512, overlap: int = 64) -> Dict:
"""Full large-orthomosaic pipeline, chaining the split stages (local/CLI use)."""
validate_upload(image_path)
original_image = load_robust_image(image_path, max_dim=max_display)
h, w = original_image.shape[:2]
print(f" Image loaded: {w}×{h}px")
if h <= tile_size and w <= tile_size:
resized = cv2.resize(original_image, (tile_size, tile_size))
raw = self.run_model(feature, resized.astype(np.float32) / 255.0)
raw_full = cv2.resize(raw, (w, h))
return self.postprocess(original_image, raw_full, feature, threshold)
tiles = self.prepare_tiles(original_image, tile_size, overlap)
preds = self.run_model_tiles(feature, tiles["chips"])
mask = self.stitch_tiles(preds, tiles["coords"], tiles["weight_map"],
tiles["image_shape"], threshold)
overlay = self._create_overlay(original_image, mask)
detected_pixels = int(np.sum(mask > 0))
total_pixels = int(mask.size)
detection_ratio = detected_pixels / total_pixels if total_pixels > 0 else 0
rooftop_analysis = None
if feature == "buildings" and detected_pixels > 0:
rooftop_analysis = self.classify_rooftops(original_image, mask)
return {
"original_image": original_image,
"mask": mask,
"overlay": overlay,
"feature": feature,
"detected_ratio": detection_ratio,
"detected_pixels": detected_pixels,
"total_pixels": total_pixels,
"rooftop_analysis": rooftop_analysis,
}
def batch_predict(self, image_paths: list, feature: str = "buildings") -> list:
results = []
for image_path in image_paths:
try:
results.append(self.predict_large(image_path, feature))
except Exception as e:
print(f"Failed to process {image_path}: {e}")
results.append(None)
return results
# Backward compatibility aliases
GramDrishtiSegmentationModel = AerixSegmentationModel
VaayuSegmentationModel = AerixSegmentationModel
def main():
import argparse
parser = argparse.ArgumentParser(description="Aerix Segmentation Inference")
parser.add_argument("image_path", help="Path to input image")
parser.add_argument("--feature", choices=["buildings", "roads", "water_bodies"], default="buildings")
parser.add_argument("--threshold", type=float, default=0.5)
parser.add_argument("--output_dir", default="inference_outputs")
args = parser.parse_args()
model = AerixSegmentationModel()
result = model.predict_large(args.image_path, args.feature, args.threshold)
output_dir = Path(args.output_dir)
output_dir.mkdir(exist_ok=True)
mask_path = output_dir / f"{Path(args.image_path).stem}_mask.png"
cv2.imwrite(str(mask_path), result["mask"])
overlay_path = output_dir / f"{Path(args.image_path).stem}_overlay.png"
overlay_bgr = cv2.cvtColor(result["overlay"], cv2.COLOR_RGB2BGR)
cv2.imwrite(str(overlay_path), overlay_bgr)
print(f"\n✓ Inference complete for {args.feature}")
print(f" Detection ratio: {result['detected_ratio']:.2%}")
print(f" Detected pixels: {result['detected_pixels']}/{result['total_pixels']}")
print(f" Mask saved: {mask_path}")
print(f" Overlay saved: {overlay_path}")
if __name__ == "__main__":
main()