Edge001's picture
Upload folder using huggingface_hub
c2a61b6 verified
Raw
History Blame Contribute Delete
17.9 kB
"""
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")