""" 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()