satellite / model_router.py
prateeksharmacoder's picture
Deploy ZeroGPU compatible code with Sen2SR
8d928e8 verified
Raw History Blame Contribute Delete
5.82 kB
import gc
import torch
from pathlib import Path
class SatelliteSRRouter:
def __init__(self, device=None):
from models.hatsat.hatsat_inference import HATSATInference
from models.esrgan.esrgan_inference import ESRGANInference
from models.sen2sr.sen2sr_inference import Sen2SRInference
self.HATSATInference = HATSATInference
self.ESRGANInference = ESRGANInference
self.Sen2SRInference = Sen2SRInference
if device is None:
device = "cuda" if torch.cuda.is_available() else "cpu"
self.device = device
# Lazy loading:
# Only the selected model is loaded.
self.hatsat = None
self.esrgan = None
self.sen2sr = None
print("=" * 55)
print("SATELLITE SR ROUTER")
print("=" * 55)
print("Device:", self.device)
print("Router initialized.")
print("Models will be loaded only when selected.")
def _clear_gpu(self):
gc.collect()
if torch.cuda.is_available():
torch.cuda.empty_cache()
torch.cuda.ipc_collect()
def _unload_hatsat(self):
if self.hatsat is not None:
print("Unloading HATSAT...")
try:
if hasattr(self.hatsat, "model"):
self.hatsat.model.cpu()
except Exception:
pass
self.hatsat = None
self._clear_gpu()
def _unload_esrgan(self):
if self.esrgan is not None:
print("Unloading ESRGAN...")
try:
if hasattr(self.esrgan, "model"):
self.esrgan.model.cpu()
except Exception:
pass
self.esrgan = None
self._clear_gpu()
def _unload_sen2sr(self):
if self.sen2sr is not None:
print("Unloading Sen2SR...")
try:
if hasattr(self.sen2sr, "model") and hasattr(self.sen2sr.model, "model"):
self.sen2sr.model.model.cpu()
except Exception:
pass
self.sen2sr = None
self._clear_gpu()
def _load_hatsat(self):
if self.hatsat is None:
print("Loading HATSAT...")
self.hatsat = self.HATSATInference(
device=self.device
)
print("HATSAT loaded successfully.")
def _load_esrgan(self):
if self.esrgan is None:
print("Loading ESRGAN...")
checkpoint = (
Path(__file__).resolve().parent
/ "weights"
/ "esrgan"
/ "RRDB_ESRGAN_x4.pth"
)
self.esrgan = self.ESRGANInference(
checkpoint_path=checkpoint,
device=self.device
)
print("ESRGAN loaded successfully.")
def _load_sen2sr(self):
if self.sen2sr is None:
print("Loading Sen2SR...")
self.sen2sr = self.Sen2SRInference(
device=self.device
)
print("Sen2SR loaded successfully.")
def available_models(self):
return {
"hatsat": {
"name": "HATSAT",
"description": "Satellite-oriented super-resolution model",
"scale": 4
},
"esrgan": {
"name": "ESRGAN",
"description": "General-purpose super-resolution baseline",
"scale": 4
},
"sen2sr": {
"name": "Sen2SR",
"description": "WEO-SAS Sentinel-2 CNN super-resolution (HuggingFace)",
"scale": 4
}
}
def predict(self, image, model_name="hatsat"):
if image is None:
raise ValueError("Please upload an image.")
model_name = str(model_name).lower().strip()
print("=" * 55)
print("ROUTER INFERENCE")
print("Selected model:", model_name)
print("=" * 55)
# ==========================================
# HATSAT
# ==========================================
if model_name == "hatsat":
# Free GPU memory used by other models
self._unload_esrgan()
self._unload_sen2sr()
# Load HATSAT only when required
self._load_hatsat()
print("Running HATSAT...")
result = self.hatsat.predict(image)
print("HATSAT inference complete.")
return result
# ==========================================
# ESRGAN
# ==========================================
elif model_name == "esrgan":
# Free GPU memory used by HATSAT and Sen2SR
self._unload_hatsat()
self._unload_sen2sr()
# Load ESRGAN only when required
self._load_esrgan()
print("Running ESRGAN...")
result = self.esrgan.predict(image)
print("ESRGAN inference complete.")
return result
# ==========================================
# SEN2SR
# ==========================================
elif model_name == "sen2sr":
# Free GPU memory used by other models
self._unload_hatsat()
self._unload_esrgan()
# Load Sen2SR only when required
self._load_sen2sr()
print("Running Sen2SR...")
result = self.sen2sr.predict(image)
print("Sen2SR inference complete.")
return result
else:
raise ValueError(
f"Unknown model '{model_name}'. "
f"Choose 'hatsat', 'esrgan', or 'sen2sr'."
)