Spaces:
Sleeping
Sleeping
Download inference.py from Rudraaaa76/geospatial-api: direct link, hf CLI and curl.
- Browser
- Download file 24.4 kB
-
https://huggingface.co/spaces/Rudraaaa76/geospatial-api/resolve/main/inference.py
- Command line
-
hf download hf://spaces/Rudraaaa76/geospatial-api/inference.py
-
curl -L -o inference.py https://huggingface.co/spaces/Rudraaaa76/geospatial-api/resolve/main/inference.py
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) | |
| 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() |