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'." )