| """ |
| 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 |
| _cache = PredictionCache() |
| LAT_MIN, LAT_MAX = 6.0, 38.0 |
| LON_MIN, LON_MAX = 68.0, 98.0 |
| |
| |
| |
| _CMAP_VMIN = 258.15 |
| _CMAP_VMAX = 308.15 |
|
|
|
|
| @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") |
|
|
| |
| _physics = PhysicsPreprocessor(_data) |
|
|
| |
| 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)") |
|
|
| |
| mean_K = _data.mean(axis=0) * _std + _mean |
| _elev = ElevationCorrector(mean_K, output_shape=(H*4, W*4)) |
|
|
| |
| 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 |
|
|
| |
| 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] |
|
|
| |
| 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) |
|
|
| |
| if _elev: |
| pred_K = _elev.apply(pred_K) |
|
|
| |
| 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) |
| |
| |
| |
| 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) |
| |
| 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") |
|
|