""" FastAPI Backend — AI Climate Downscaling System Model : Swin2SR Ensemble (realworld + classical), 12.1M x2 INT8, 24.2M total Data : ERA5 2020, India, 2m Temperature Run : python serve.py | http://127.0.0.1:8080 Four improvements over baseline Swin2SR: 1. Ensemble — realworld + classical averaged -> cancels noise 2. Physics — anomaly-based inference -> model sees deviations not absolutes 3. Elevation — 6.5K/km lapse rate post-correction -> terrain accuracy 4. Land mask — AI only applied to land, ERA5 SST preserved over ocean """ import sys, io, os, time import torch, numpy as np import matplotlib; matplotlib.use("Agg") import matplotlib.pyplot as plt from PIL import Image from functools import lru_cache from contextlib import asynccontextmanager from fastapi import FastAPI, Response from fastapi.middleware.cors import CORSMiddleware from fastapi.staticfiles import StaticFiles from pathlib import Path os.environ.setdefault("PYTHONIOENCODING", "utf-8") ROOT = Path(__file__).parent.parent sys.path.insert(0, str(ROOT)) from src.utils.config import load_config from src.data.netcdf_loader import load_climate_data from src.data.preprocessor import Preprocessor from src.models.downscaler import ( load_model, PredictionCache, PhysicsPreprocessor, ElevationCorrector ) from src.models.land_mask import build_land_mask, apply_land_mask _model = None; _data = None; _preproc = None _mean = 0.0; _std = 1.0 _t_min = -3.0; _t_max = 3.0 _physics = None; _elev = None; _mask_lr = None; _mask_hr = None _mean_field_hr = None # physics mean field upsampled to HR once at startup _cache = PredictionCache() LAT_MIN, LAT_MAX = 6.0, 38.0 LON_MIN, LON_MAX = 68.0, 98.0 # Fixed colormap range calibrated to India's annual temperature spread. # Blue = <-15C (Himalayas), cyan/green = 0-15C (North India plains), # yellow/orange = 15-25C (Deccan), red = >25C (South coast + ocean) _CMAP_VMIN = 258.15 # K = -15 C _CMAP_VMAX = 308.15 # K = +35 C @asynccontextmanager async def lifespan(app: FastAPI): global _model, _data, _preproc, _mean, _std, _t_min, _t_max global _physics, _elev, _mask_lr, _mask_hr, _mean_field_hr print("\n" + "="*56) print(" ATMOS — AI Climate Downscaling Startup") print("="*56) _model = load_model(device="cpu") _model.eval() nc = ROOT / "data" / "raw" / "era5_real_2020-01-01_2020-12-31.nc" if nc.exists(): cfg = load_config() loader = load_climate_data(cfg, data_path=str(nc)) raw = loader.get_variable(cfg["data"]["variables"][0]) loader.close() _preproc = Preprocessor(cfg) _data, _ = _preproc.fit_transform(raw, cfg["data"]["variables"][0]) stats = _preproc.statistics[cfg["data"]["variables"][0]] _mean = float(stats["mean"]) _std = float(stats["std"]) _t_min = float(np.percentile(_data, 1)) _t_max = float(np.percentile(_data, 99)) T, H, W = _data.shape print(f" ERA5: {T} timesteps | {H}x{W} | mean={_mean:.1f}K") # IMPROVEMENT 2: physics pre-processor _physics = PhysicsPreprocessor(_data) # Precompute physics mean field at HR resolution — avoids scipy zoom per inference from scipy.ndimage import zoom as spz _mean_field_hr = spz(_physics.mean_field, 4, order=3)[:H*4, :W*4] print(f" [Physics] Mean field upsampled to {_mean_field_hr.shape} (precomputed)") # IMPROVEMENT 3: elevation corrector (needs mean in Kelvin) mean_K = _data.mean(axis=0) * _std + _mean _elev = ElevationCorrector(mean_K, output_shape=(H*4, W*4)) # IMPROVEMENT 4: land-sea mask print(" [Mask] Building land-sea mask from temporal variance...") _mask_lr, _mask_hr = build_land_mask( _data, LAT_MIN, LAT_MAX, LON_MIN, LON_MAX, output_scale=4) print(" [Cache] POST /api/cache/build to start background pre-compute") else: print(f" WARNING: ERA5 not found at {nc}") print(f" Ready -> http://127.0.0.1:8080") print("="*56 + "\n") yield _cache.stop() print("Shutdown.") app = FastAPI(title="ATMOS — Swin2SR Ensemble", lifespan=lifespan) app.add_middleware(CORSMiddleware, allow_origins=["*"], allow_methods=["*"], allow_headers=["*"]) def _no_data(): return Response(content='{"error":"data not loaded"}', status_code=503, media_type="application/json") def _z2k(z): return z * _std + _mean def _run_inference(t: int): """ Full pipeline: physics pre-process -> ensemble inference -> elevation fix. Serves from cache when available (instant), live inference otherwise. """ era5 = _data[t].copy() era5_K = _z2k(era5) cached = _cache.get(t) if cached is not None: return era5_K, cached.astype(np.float32) H, W = era5.shape ph = (64 - H % 64) % 64 pw = (64 - W % 64) % 64 # IMPROVEMENT 2: anomaly pre-processing inp = _physics.to_anomaly(era5) if _physics else era5 padded = np.pad(inp, ((0, ph), (0, pw)), mode="edge") x = torch.from_numpy(padded).unsqueeze(0).unsqueeze(0).float() with torch.no_grad(): out = _model(x).squeeze().numpy() pred = out[:H*4, :W*4] # IMPROVEMENT 2: add mean back (use precomputed HR mean field) if _physics and _mean_field_hr is not None: pred = _physics.from_anomaly(pred, _mean_field_hr) elif _physics: from scipy.ndimage import zoom as spz pred = _physics.from_anomaly(pred, spz(_physics.mean_field, 4, order=3)[:H*4, :W*4]) pred_K = _z2k(pred) # IMPROVEMENT 3: elevation correction if _elev: pred_K = _elev.apply(pred_K) # IMPROVEMENT 4: land-sea mask — preserve ERA5 SST over ocean if _mask_hr is not None: pred_K = apply_land_mask(pred_K, era5_K, _mask_hr) return era5_K, pred_K @lru_cache(maxsize=16) def _cached_inference(t: int): return _run_inference(t) def _to_png(arr_K, vmin=None, vmax=None, cmap="RdYlBu_r", alpha=220, sharpen=False): """Render temperature array as PNG. sharpen=True for AI output only.""" data = arr_K.copy() if sharpen: from scipy.ndimage import gaussian_filter blurred = gaussian_filter(data, sigma=1.2) data = data + 1.8 * (data - blurred) v0 = vmin if vmin is not None else float(data.min()) v1 = vmax if vmax is not None else float(data.max()) norm = np.clip((data - v0) / (v1 - v0 + 1e-8), 0, 1) rgba = (plt.get_cmap(cmap)(norm) * 255).astype(np.uint8) rgba[..., 3] = alpha buf = io.BytesIO() Image.fromarray(rgba, "RGBA").save(buf, format="PNG", compress_level=1) buf.seek(0) return Response(content=buf.getvalue(), media_type="image/png") def _err_to_png(pred_K, era5_K): from scipy.ndimage import zoom as spz f = pred_K.shape[0] / era5_K.shape[0] up = spz(era5_K, f, order=1)[:pred_K.shape[0], :pred_K.shape[1]] err = pred_K - up ab = max(abs(err.min()), abs(err.max())) + 1e-4 norm = np.clip((err + ab) / (2*ab), 0, 1) rgba = (plt.get_cmap("RdBu_r")(norm) * 255).astype(np.uint8) rgba[..., 3] = 200 buf = io.BytesIO() Image.fromarray(rgba, "RGBA").save(buf, format="PNG", compress_level=1) buf.seek(0) return Response(content=buf.getvalue(), media_type="image/png") def _laplacian(arr): gy, gx = np.gradient(arr) gyy, _ = np.gradient(gy); _, gxx = np.gradient(gx) return float(np.mean(np.abs(gyy + gxx))) def _sharpness_gain(era5_K, pred_K): from scipy.ndimage import zoom as spz f = pred_K.shape[0] / era5_K.shape[0] up = spz(era5_K, f, order=3)[:pred_K.shape[0], :pred_K.shape[1]] s0 = _laplacian(up); s1 = _laplacian(pred_K) return s0, s1, round((s1 / (s0 + 1e-8) - 1.0) * 100, 1) @app.get("/api/health") def health(): return {"status": "ok", "model_loaded": _model is not None, "data_loaded": _data is not None, "improvements": ["ensemble", "physics_anomaly", "elevation_correction", "land_sea_mask"], "timesteps": int(_data.shape[0]) if _data is not None else 0} @app.get("/api/meta") def meta(): if _data is None: return {"timesteps": 0, "error": "data not loaded"} T, H, W = _data.shape import psutil; vm = psutil.virtual_memory() return { "timesteps": T, "grid_shape": [H, W], "output_shape": [H*4, W*4], "bounds": [[LAT_MIN, LON_MIN], [LAT_MAX, LON_MAX]], "model": "Swin2SR Ensemble (realworld + classical)", "params_m": 24.2, "model_ids": [ "caidas/swin2SR-realworld-sr-x4-64-bsrgan-psnr", "caidas/swin2SR-classical-sr-x4-64" ], "quantisation": "INT8 dynamic (both models)", "improvements": { "ensemble": "realworld + classical averaged", "physics": "anomaly-based inference (temporal mean removed)", "elevation": "6.5K/km lapse rate correction", "land_mask": "ocean pixels use ERA5 SST, AI only applied to land" }, "threads": torch.get_num_threads(), "region": "India (6N-38N, 68E-98E)", "variable": "2m Temperature", "source": "ERA5 Reanalysis 2020", "scale_factor": 4, "input_res": "0.25deg (28km)", "output_res": "0.0625deg (7km)", "mean_K": _mean, "std_K": _std, "t_min": float(_t_min), "t_max": float(_t_max), "cache_progress": round(_cache.progress * 100, 1), "cache_ready": _cache.ready, "ram_pct": round(vm.percent, 1), "ram_used_gb": round(vm.used / 1e9, 1) } @app.get("/api/map/{layer}/{t_idx}.png") def map_image(layer: str, t_idx: int): if _data is None: return _no_data() if not (0 <= t_idx < _data.shape[0]): return Response(status_code=404) era5_K, pred_K = _cached_inference(t_idx) # era5 at LR resolution (129x121) uses mask_lr # pred at HR resolution (516x484) uses mask_hr # sharpen=True only on AI output — makes improvement visually obvious if layer == "lr": return _to_png(era5_K) elif layer == "pred": return _to_png(pred_K, sharpen=True) elif layer == "hr": return _to_png(era5_K) elif layer == "error": return _err_to_png(pred_K, era5_K) return Response(status_code=400) @app.get("/api/stats/{t_idx}") def stats(t_idx: int): if _data is None: return _no_data() if not (0 <= t_idx < _data.shape[0]): return Response(status_code=404) era5_K, pred_K = _cached_inference(t_idx) s0, s1, sg = _sharpness_gain(era5_K, pred_K) detail = float(np.std(pred_K) / (np.std(era5_K) + 1e-6)) import psutil; vm = psutil.virtual_memory() return { "mean_temp_K": round(float(np.mean(era5_K)), 2), "mean_temp_C": round(float(np.mean(era5_K)) - 273.15, 1), "min_temp_K": round(float(era5_K.min()), 2), "max_temp_K": round(float(era5_K.max()), 2), "std_era5": round(float(np.std(era5_K)), 4), "std_pred": round(float(np.std(pred_K)), 4), "spatial_detail_gain": round(detail, 3), "sharpness_era5": round(s0, 4), "sharpness_pred": round(s1, 4), "sharpness_gain_pct": sg, "cache_progress": round(_cache.progress * 100, 1), "cache_ready": _cache.ready, "ram_pct": round(vm.percent, 1), "model": "Swin2SR Ensemble + Physics + Elevation" } @app.get("/api/point/{t_idx}/{lat}/{lon}") def point_value(t_idx: int, lat: float, lon: float): if _data is None: return _no_data() era5_K, pred_K = _cached_inference(t_idx) def rc(arr): H, W = arr.shape r = max(0, min(H-1, int((LAT_MAX - lat) / (LAT_MAX - LAT_MIN) * H))) c = max(0, min(W-1, int((lon - LON_MIN) / (LON_MAX - LON_MIN) * W))) return r, c r0, c0 = rc(era5_K); r1, c1 = rc(pred_K) return {"lat": lat, "lon": lon, "era5_K": round(float(era5_K[r0, c0]), 2), "pred_K": round(float(pred_K[r1, c1]), 2), "era5_C": round(float(era5_K[r0, c0]) - 273.15, 1), "pred_C": round(float(pred_K[r1, c1]) - 273.15, 1)} @app.get("/api/timeseries/{lat}/{lon}") def timeseries(lat: float, lon: float, step: int = 24): if _data is None: return _no_data() T, H, W = _data.shape r = max(0, min(H-1, int((LAT_MAX - lat) / (LAT_MAX - LAT_MIN) * H))) c = max(0, min(W-1, int((lon - LON_MIN) / (LON_MAX - LON_MIN) * W))) idxs = list(range(0, T, step)) vals_K = [round(float(_data[i, r, c]) * _std + _mean, 2) for i in idxs] return {"indices": idxs, "values_K": vals_K, "values_C": [round(v - 273.15, 1) for v in vals_K], "lat": lat, "lon": lon} @lru_cache(maxsize=6) def _compute_uncertainty(t: int): era5 = _data[t].copy(); H, W = era5.shape ph = (64 - H % 64) % 64; pw = (64 - W % 64) % 64 jitters = [(0,0),(1,0),(-1,0),(0,1),(0,-1),(1,1),(-1,1),(1,-1)] preds = [] for dy, dx in jitters: sh = np.roll(np.roll(era5, dy, 0), dx, 1) if dy > 0: sh[:dy, :] = era5[:dy, :] elif dy < 0: sh[dy:, :] = era5[dy:, :] if dx > 0: sh[:, :dx] = era5[:, :dx] elif dx < 0: sh[:, dx:] = era5[:, dx:] inp = _physics.to_anomaly(sh) if _physics else sh padded = np.pad(inp, ((0, ph), (0, pw)), mode="edge") x = torch.from_numpy(padded).unsqueeze(0).unsqueeze(0).float() with torch.no_grad(): out = _model(x).squeeze().numpy() pred = out[:H*4, :W*4] if _physics and _mean_field_hr is not None: pred = _physics.from_anomaly(pred, _mean_field_hr) elif _physics: from scipy.ndimage import zoom as spz pred = _physics.from_anomaly(pred, spz(_physics.mean_field, 4, order=3)[:H*4, :W*4]) pred_K = _z2k(pred) # Apply land mask — ocean uncertainty = 0 (ERA5 passthrough, no AI variance) if _mask_hr is not None: from scipy.ndimage import zoom as spz era5_K_unc = _z2k(era5) era5_up = spz(era5_K_unc, 4, order=1)[:H*4, :W*4] pred_K = np.where(_mask_hr[:H*4,:W*4], pred_K, era5_up) pred = np.roll(np.roll(pred_K, -dy*4, 0), -dx*4, 1) preds.append(pred) stack = np.stack(preds, 0) return np.std(stack, 0), np.mean(stack, 0) @lru_cache(maxsize=6) def _compute_psd(t: int): from scipy.ndimage import zoom as spz era5_K, pred_K = _cached_inference(t) f = pred_K.shape[0] / era5_K.shape[0] up = spz(era5_K, f, order=3)[:pred_K.shape[0], :pred_K.shape[1]] def rpsd(arr): arr = arr - arr.mean(); H, W = arr.shape arr = arr * np.hanning(H)[:, None] * np.hanning(W)[None, :] F = np.fft.fftshift(np.fft.fft2(arr)); P = (np.abs(F)**2) / (H*W) cy, cx = H//2, W//2 Y, X = np.mgrid[-cy:H-cy, -cx:W-cx] R = np.sqrt(X**2 + Y**2).astype(int); mr = min(cy, cx) return np.array([P[R==r].mean() if np.any(R==r) else 0 for r in range(1, mr)]) p0 = rpsd(up); p1 = rpsd(pred_K); n = min(len(p0), len(p1)) wl = (min(pred_K.shape) / np.arange(1, n+1)) * 7.0 mask = (wl >= 10) & (wl <= 500) return {"wavelengths_km": wl[mask].tolist(), "psd_era5": np.log10(p0[:n][mask] + 1e-20).tolist(), "psd_ai": np.log10(p1[:n][mask] + 1e-20).tolist()} @app.get("/api/uncertainty/{t_idx}.png") def uncertainty_map(t_idx: int): if _data is None: return _no_data() if not (0 <= t_idx < _data.shape[0]): return Response(status_code=404) unc, _ = _compute_uncertainty(t_idx) vmax = float(np.percentile(unc, 98)) return _to_png(unc, vmin=0.0, vmax=max(vmax, 0.1), cmap="hot", alpha=210) @app.get("/api/uncertainty/{t_idx}/stats") def uncertainty_stats(t_idx: int): if _data is None: return _no_data() unc, _ = _compute_uncertainty(t_idx) return {"mean_uncertainty_K": round(float(unc.mean()), 4), "max_uncertainty_K": round(float(unc.max()), 4), "p95_uncertainty_K": round(float(np.percentile(unc, 95)), 4), "high_uncertainty_pct": round(float((unc > unc.mean()+unc.std()).mean()*100), 1), "method": "TTA Spatial Jitter (8 directions)", "n_samples": 8} @app.get("/api/psd/{t_idx}") def psd(t_idx: int): if _data is None: return _no_data() if not (0 <= t_idx < _data.shape[0]): return Response(status_code=404) return _compute_psd(t_idx) @app.get("/api/cache/status") def cache_status(): import psutil; vm = psutil.virtual_memory() T = _data.shape[0] if _data is not None else 0 return {"cached_frames": len(_cache._cache), "total_frames": T, "progress_pct": round(_cache.progress * 100, 1), "ready": _cache.ready, "eta_sec": _cache.eta_sec, "ram_used_gb": round(vm.used/1e9, 1), "ram_pct": round(vm.percent, 1)} @app.post("/api/cache/build") def cache_build(workers: int = 2): if _data is None: return _no_data() if _cache.progress > 0 and not _cache.ready: return {"status": "already_running", "progress_pct": round(_cache.progress*100, 1)} if _cache.ready: return {"status": "already_complete", "cached_frames": len(_cache._cache)} _cache.start( model = _model, data = _data, zscore_fn= lambda z: z * _std + _mean, physics = _physics, elevation = _elev, mask_hr = _mask_hr, workers = max(1, min(workers, 6)), on_progress = lambda d, t, e: print(f" [Cache] {d}/{t} ({d/t*100:.1f}%) " f"ETA {e//60}m{e%60:02d}s", end="\r") ) return {"status": "started", "workers": workers, "total_frames": int(_data.shape[0])} app.mount("/", StaticFiles(directory=str(ROOT/"dashboard_frontend"), html=True), name="frontend")