Spaces:
Running on Zero
Running on Zero
Download model_router.py from prateeksharmacoder/satellite: direct link, hf CLI and curl.
- Browser
- Download file 5.82 kB
-
https://huggingface.co/spaces/prateeksharmacoder/satellite/resolve/main/model_router.py
- Command line
-
hf download hf://spaces/prateeksharmacoder/satellite/model_router.py
-
curl -L -o model_router.py https://huggingface.co/spaces/prateeksharmacoder/satellite/resolve/main/model_router.py
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'." | |
| ) | |