Spaces:
Running on Zero
Running on Zero
| """ | |
| Phase 5 -- Scientific Validation (PSNR / SSIM / RMSE) | |
| satellite_sr_deploy/validation.py | |
| We use the standard Super-Resolution evaluation protocol since true | |
| higher-resolution ground truth (e.g. 2.5m WorldView) is unavailable | |
| for this specific Sentinel-2 scene. | |
| Protocol: | |
| 1. Load original 10m Sentinel-2 image (Ground Truth). | |
| 2. Crop to dimensions divisible by 4. | |
| 3. Downsample by 4x using Bicubic interpolation (simulating 40m input). | |
| 4. Upsample the 40m image back to 10m using Bicubic, ESRGAN, and HATSAT. | |
| 5. Compute PSNR, SSIM, and RMSE between the upsampled images and the GT. | |
| """ | |
| import sys | |
| import math | |
| import time | |
| from pathlib import Path | |
| import torch | |
| import numpy as np | |
| from PIL import Image | |
| # Ensure satellite_sr_deploy is in path | |
| ROOT = Path(__file__).resolve().parent | |
| if str(ROOT) not in sys.path: | |
| sys.path.insert(0, str(ROOT)) | |
| from model_router import SatelliteSRRouter | |
| def calculate_psnr(img1, img2): | |
| """Calculate Peak Signal-to-Noise Ratio (PSNR) in RGB space.""" | |
| # img1 and img2 can be numpy arrays | |
| img1 = img1.astype(np.float64) | |
| img2 = img2.astype(np.float64) | |
| mse = np.mean((img1 - img2) ** 2) | |
| if mse == 0: | |
| return float('inf') | |
| return 20 * math.log10(255.0 / math.sqrt(mse)) | |
| def calculate_rmse(img1, img2): | |
| """Calculate Root Mean Square Error (RMSE).""" | |
| img1 = img1.astype(np.float64) | |
| img2 = img2.astype(np.float64) | |
| return np.sqrt(np.mean((img1 - img2) ** 2)) | |
| def ssim_core(img1, img2, C1=6.5025, C2=58.5225): | |
| """ | |
| Simplified SSIM over the entire image. | |
| Standard SSIM uses a gaussian window, but for a quick numpy implementation | |
| we can approximate or compute global/block-based SSIM. | |
| For robust SSIM without adding scikit-image dependency, we'll compute | |
| global mean/variance which gives a rough structural similarity. | |
| (Note: True SSIM uses local sliding windows, so we'll implement a simple | |
| 11x11 sliding window in numpy if possible, or just use the global version | |
| for demonstration since this is a local metric). | |
| """ | |
| # For strict SSIM, we should ideally use torchvision or skimage. | |
| # Let's see if skimage is available. | |
| pass | |
| try: | |
| from skimage.metrics import structural_similarity as ssim | |
| HAS_SKIMAGE = True | |
| except ImportError: | |
| HAS_SKIMAGE = False | |
| print("skimage not found. SSIM will be calculated using a global approximation.") | |
| def calculate_ssim(img1, img2): | |
| if HAS_SKIMAGE: | |
| # skimage expects channel_axis=-1 for RGB | |
| return ssim(img1, img2, channel_axis=-1, data_range=255) | |
| else: | |
| # Global approximation (Not strictly standard SSIM, but gives a relative metric) | |
| img1 = img1.astype(np.float64) | |
| img2 = img2.astype(np.float64) | |
| mu1 = img1.mean() | |
| mu2 = img2.mean() | |
| sigma1 = img1.var() | |
| sigma2 = img2.var() | |
| cov = np.cov(img1.flatten(), img2.flatten())[0, 1] | |
| c1 = (0.01 * 255)**2 | |
| c2 = (0.03 * 255)**2 | |
| num = (2 * mu1 * mu2 + c1) * (2 * cov + c2) | |
| den = (mu1**2 + mu2**2 + c1) * (sigma1 + sigma2 + c2) | |
| return num / den | |
| def run_validation(): | |
| print("=" * 60) | |
| print("PHASE 5: SCIENTIFIC VALIDATION") | |
| print("=" * 60) | |
| gt_path = ROOT.parent / "data" / "processed" / "sentinel2_rgb_10m.png" | |
| if not gt_path.exists(): | |
| print(f"Error: Could not find GT image at {gt_path}") | |
| sys.exit(1) | |
| gt_img = Image.open(gt_path).convert("RGB") | |
| orig_w, orig_h = gt_img.size | |
| print(f"Original Ground Truth (GT) size: {orig_w}x{orig_h}") | |
| # 1. Crop GT to a multiple of 4 | |
| w = (orig_w // 4) * 4 | |
| h = (orig_h // 4) * 4 | |
| gt_img = gt_img.crop((0, 0, w, h)) | |
| print(f"Cropped GT size (divisible by 4): {w}x{h}") | |
| # 2. Downsample by 4x to simulate 40m input | |
| lr_w, lr_h = w // 4, h // 4 | |
| lr_img = gt_img.resize((lr_w, lr_h), Image.BICUBIC) | |
| print(f"Low-Resolution (LR) 40m simulated size: {lr_w}x{lr_h}") | |
| # 3. Upsample using Bicubic | |
| print("\n--- Running Baseline: Bicubic Interpolation ---") | |
| start_time = time.time() | |
| bicubic_img = lr_img.resize((w, h), Image.BICUBIC) | |
| bicubic_time = time.time() - start_time | |
| # Initialize Models | |
| print("\n--- Initializing Deep Learning Models ---") | |
| router = SatelliteSRRouter() | |
| # 4. Upsample using ESRGAN | |
| print("\n--- Running ESRGAN ---") | |
| start_time = time.time() | |
| esrgan_img = router.predict(lr_img, "esrgan") | |
| esrgan_time = time.time() - start_time | |
| # 5. Upsample using HATSAT | |
| print("\n--- Running HATSAT ---") | |
| start_time = time.time() | |
| hatsat_img = router.predict(lr_img, "hatsat") | |
| hatsat_time = time.time() - start_time | |
| # 6. Evaluation | |
| print("\n" + "=" * 60) | |
| print("EVALUATION RESULTS (vs Ground Truth)") | |
| print("=" * 60) | |
| gt_np = np.array(gt_img) | |
| bicubic_np = np.array(bicubic_img) | |
| esrgan_np = np.array(esrgan_img) | |
| hatsat_np = np.array(hatsat_img) | |
| results = {} | |
| for name, img_np, t in [("Bicubic", bicubic_np, bicubic_time), | |
| ("ESRGAN", esrgan_np, esrgan_time), | |
| ("HATSAT", hatsat_np, hatsat_time)]: | |
| # Ensure dimensions match (in case of any cropping/padding discrepancy) | |
| if img_np.shape != gt_np.shape: | |
| # Crop to match GT exactly | |
| img_np = img_np[:h, :w, :] | |
| psnr_val = calculate_psnr(gt_np, img_np) | |
| ssim_val = calculate_ssim(gt_np, img_np) | |
| rmse_val = calculate_rmse(gt_np, img_np) | |
| results[name] = { | |
| "PSNR": psnr_val, | |
| "SSIM": ssim_val, | |
| "RMSE": rmse_val, | |
| "Time": t | |
| } | |
| print(f"{name.upper()}:") | |
| print(f" PSNR: {psnr_val:.2f} dB (higher is better)") | |
| print(f" SSIM: {ssim_val:.4f} (higher is better)") | |
| print(f" RMSE: {rmse_val:.2f} (lower is better)") | |
| print(f" Time: {t:.2f} s") | |
| print("-" * 40) | |
| print("SUMMARY CONCLUSION:") | |
| best_psnr = max(results, key=lambda k: results[k]["PSNR"]) | |
| best_ssim = max(results, key=lambda k: results[k]["SSIM"]) | |
| best_rmse = min(results, key=lambda k: results[k]["RMSE"]) | |
| print(f"Highest PSNR: {best_psnr} ({results[best_psnr]['PSNR']:.2f} dB)") | |
| print(f"Highest SSIM: {best_ssim} ({results[best_ssim]['SSIM']:.4f})") | |
| print(f"Lowest RMSE: {best_rmse} ({results[best_rmse]['RMSE']:.2f})") | |
| # Save the synthetic LR and upsampled variants for visual inspection | |
| out_dir = ROOT / "outputs" / "validation" | |
| out_dir.mkdir(parents=True, exist_ok=True) | |
| gt_img.save(out_dir / "01_GroundTruth.png") | |
| lr_img.save(out_dir / "02_SimulatedLR.png") | |
| bicubic_img.save(out_dir / "03_Bicubic.png") | |
| esrgan_img.save(out_dir / "04_ESRGAN.png") | |
| hatsat_img.save(out_dir / "05_HATSAT.png") | |
| print(f"\nSaved visual comparison images to: {out_dir}") | |
| print("=" * 60) | |
| if __name__ == "__main__": | |
| run_validation() | |