satellite / local_test.py
prateeksharmacoder's picture
Deploy ZeroGPU compatible code with Sen2SR
8d928e8 verified
Raw
History Blame Contribute Delete
6.31 kB
"""
Local test script for satellite_sr_deploy.
Tests:
1. Import checks for all modules
2. HATSAT inference on a synthetic 198x139 image
3. ESRGAN inference on a synthetic 198x139 image
4. Router switch (HATSAT β†’ ESRGAN)
"""
import sys
import time
from pathlib import Path
# -------------------------------------------------------
# Make sure we run from the deploy root
# -------------------------------------------------------
ROOT = Path(__file__).resolve().parent
if str(ROOT) not in sys.path:
sys.path.insert(0, str(ROOT))
print("=" * 60)
print("STEP 1 β€” IMPORT CHECK")
print("=" * 60)
try:
import torch
print(f" torch : {torch.__version__}")
except ImportError as e:
print(f" [FAIL] torch : {e}")
sys.exit(1)
try:
import numpy as np
print(f" numpy : {np.__version__}")
except ImportError as e:
print(f" [FAIL] numpy : {e}")
sys.exit(1)
try:
from PIL import Image
import PIL
print(f" Pillow : {PIL.__version__}")
except ImportError as e:
print(f" [FAIL] Pillow : {e}")
sys.exit(1)
try:
import einops
print(f" einops : {einops.__version__}")
except ImportError as e:
print(f" [FAIL] einops : {e}")
sys.exit(1)
try:
import gradio as gr
print(f" gradio : {gr.__version__}")
except ImportError as e:
print(f" [FAIL] gradio : {e}")
sys.exit(1)
try:
import rasterio
print(f" rasterio : {rasterio.__version__}")
except ImportError as e:
print(f" [FAIL] rasterio : {e}")
sys.exit(1)
print("\nAll imports OK.\n")
# -------------------------------------------------------
# Check weight files
# -------------------------------------------------------
print("=" * 60)
print("STEP 2 β€” WEIGHT FILE CHECK")
print("=" * 60)
hatsat_ckpt = ROOT / "weights" / "hatsat" / "net_g_150000.pth"
esrgan_ckpt = ROOT / "weights" / "esrgan" / "RRDB_ESRGAN_x4.pth"
if hatsat_ckpt.exists():
size_mb = hatsat_ckpt.stat().st_size / (1024 * 1024)
print(f" HATSAT : {hatsat_ckpt} ({size_mb:.1f} MB) OK")
else:
print(f" [FAIL] HATSAT checkpoint not found: {hatsat_ckpt}")
sys.exit(1)
if esrgan_ckpt.exists():
size_mb = esrgan_ckpt.stat().st_size / (1024 * 1024)
print(f" ESRGAN : {esrgan_ckpt} ({size_mb:.1f} MB) OK")
else:
print(f" [FAIL] ESRGAN checkpoint not found: {esrgan_ckpt}")
sys.exit(1)
print()
# -------------------------------------------------------
# Model router import
# -------------------------------------------------------
print("=" * 60)
print("STEP 3 β€” MODEL ROUTER IMPORT")
print("=" * 60)
try:
from model_router import SatelliteSRRouter
print(" SatelliteSRRouter imported OK")
except Exception as e:
print(f" [FAIL] model_router import: {e}")
sys.exit(1)
print()
# -------------------------------------------------------
# Create a synthetic test image (198 x 139, RGB)
# -------------------------------------------------------
print("=" * 60)
print("STEP 4 β€” CREATE SYNTHETIC TEST IMAGE (198x139 RGB)")
print("=" * 60)
import numpy as np
test_arr = np.random.randint(0, 255, (139, 198, 3), dtype=np.uint8)
test_image = Image.fromarray(test_arr, mode="RGB")
print(f" Input size: {test_image.size} (W x H)")
print()
# -------------------------------------------------------
# Instantiate router (lazy β€” no model loaded yet)
# -------------------------------------------------------
print("=" * 60)
print("STEP 5 β€” INIT ROUTER")
print("=" * 60)
device = "cuda" if torch.cuda.is_available() else "cpu"
print(f" Device: {device}")
router = SatelliteSRRouter(device=device)
print()
# -------------------------------------------------------
# HATSAT inference
# -------------------------------------------------------
print("=" * 60)
print("STEP 6 β€” HATSAT INFERENCE")
print("=" * 60)
t0 = time.time()
try:
hatsat_output = router.predict(test_image, model_name="hatsat")
elapsed = time.time() - t0
print(f" Input : {test_image.size}")
print(f" Output : {hatsat_output.size}")
expected = (test_image.width * 4, test_image.height * 4)
if hatsat_output.size == expected:
print(f" Scale check: OK (expected {expected})")
else:
print(f" [WARN] Scale mismatch: expected {expected}, got {hatsat_output.size}")
print(f" Time : {elapsed:.2f}s")
except Exception as e:
print(f" [FAIL] HATSAT inference: {e}")
import traceback; traceback.print_exc()
sys.exit(1)
print()
# -------------------------------------------------------
# ESRGAN inference (router should unload HATSAT first)
# -------------------------------------------------------
print("=" * 60)
print("STEP 7 β€” ESRGAN INFERENCE (router switches model)")
print("=" * 60)
t0 = time.time()
try:
esrgan_output = router.predict(test_image, model_name="esrgan")
elapsed = time.time() - t0
print(f" Input : {test_image.size}")
print(f" Output : {esrgan_output.size}")
expected = (test_image.width * 4, test_image.height * 4)
if esrgan_output.size == expected:
print(f" Scale check: OK (expected {expected})")
else:
print(f" [WARN] Scale mismatch: expected {expected}, got {esrgan_output.size}")
print(f" Time : {elapsed:.2f}s")
except Exception as e:
print(f" [FAIL] ESRGAN inference: {e}")
import traceback; traceback.print_exc()
sys.exit(1)
print()
# -------------------------------------------------------
# Save outputs
# -------------------------------------------------------
print("=" * 60)
print("STEP 8 β€” SAVE TEST OUTPUTS")
print("=" * 60)
out_dir = ROOT / "outputs" / "local_test"
out_dir.mkdir(parents=True, exist_ok=True)
hatsat_path = out_dir / "test_hatsat_output.png"
esrgan_path = out_dir / "test_esrgan_output.png"
hatsat_output.save(hatsat_path)
esrgan_output.save(esrgan_path)
print(f" HATSAT output saved : {hatsat_path}")
print(f" ESRGAN output saved : {esrgan_path}")
print()
# -------------------------------------------------------
# Summary
# -------------------------------------------------------
print("=" * 60)
print("LOCAL TEST COMPLETE")
print("=" * 60)
print(" All tests passed.")
print(f" HATSAT : {test_image.size} -> {hatsat_output.size}")
print(f" ESRGAN : {test_image.size} -> {esrgan_output.size}")
print("=" * 60)