| |
| """Samudra 2 — autoregressive global ocean emulator, rolled out live. |
| |
| Streams a real initial condition from the public 1° GFDL-OM4 store, runs the |
| released `M2LInES/Samudra2` checkpoint forward autoregressively on ZeroGPU, and |
| animates the emulated ocean next to the OM4 ground truth it is emulating. |
| """ |
|
|
| from __future__ import annotations |
|
|
| import os |
| import tempfile |
| import time |
|
|
| import spaces |
|
|
| import gradio as gr |
| import matplotlib |
| import numpy as np |
| import torch |
|
|
| matplotlib.use("Agg") |
| import imageio.v2 as imageio |
| import matplotlib.pyplot as plt |
| from huggingface_hub import hf_hub_download |
|
|
| from ocean_data import HIST, N_PROG, PROG_VARS, STEP_DAYS, get_store |
| from samudra_model import build_samudra |
|
|
| try: |
| import cmocean.cm as cmo |
|
|
| CM_THERMAL, CM_HALINE, CM_BALANCE = cmo.thermal, cmo.haline, cmo.balance |
| except Exception: |
| CM_THERMAL, CM_HALINE, CM_BALANCE = "inferno", "viridis", "RdBu_r" |
|
|
| MODEL_REPO = "M2LInES/Samudra2" |
| MAX_STEPS = 36 |
|
|
| |
| print("[boot] building Samudra 2 (1deg) ...", flush=True) |
| ckpt_path = hf_hub_download(MODEL_REPO, "onedeg/ema_ckpt.pt") |
| model = build_samudra() |
| _state = torch.load(ckpt_path, map_location="cpu", weights_only=False)["model"] |
| model.load_state_dict({k.replace("module.", "", 1): v for k, v in _state.items()}) |
| model.eval().to("cuda") |
| print("[boot] model on cuda", flush=True) |
|
|
| |
| print("[boot] opening public OM4 zarr store ...", flush=True) |
| store = get_store() |
| LABEL_MASK = torch.from_numpy(np.tile(store.prog_mask, (HIST + 1, 1, 1))) |
| DATES = [store.date_str(i) for i in range(len(store.time))] |
|
|
| |
| WINDOW_START, WINDOW_END = "2014-10-10", "2022-12-24" |
| _first = next(i for i, d in enumerate(DATES) if d >= WINDOW_START) |
| _last = len(DATES) - 2 * MAX_STEPS - 2 |
| START_DATES = [DATES[i] for i in range(_first, _last + 1, 6)] |
| print(f"[boot] {len(START_DATES)} start dates, {DATES[_first]} .. {DATES[_last]}", flush=True) |
|
|
| |
| |
| VARIABLES: dict[str, tuple[str, object, str, str]] = { |
| "Sea surface temperature (0 m)": ("thetao_0", CM_THERMAL, "°C", "SST"), |
| "Sea surface height": ("zos", CM_BALANCE, "m", "SSH"), |
| "Sea surface salinity (0 m)": ("so_0", CM_HALINE, "psu", "SSS"), |
| "Temperature at 250 m": ("thetao_7", CM_THERMAL, "°C", "T 250 m"), |
| "Temperature at 1050 m": ("thetao_11", CM_THERMAL, "°C", "T 1050 m"), |
| "Zonal velocity at 0 m": ("uo_0", CM_BALANCE, "m/s", "u 0 m"), |
| "Meridional velocity at 0 m": ("vo_0", CM_BALANCE, "m/s", "v 0 m"), |
| "Salinity at 1050 m": ("so_11", CM_HALINE, "psu", "S 1050 m"), |
| } |
| DEFAULT_VARIABLE = "Sea surface temperature (0 m)" |
| DEFAULT_DATE = START_DATES[0] |
|
|
|
|
| def _gpu_duration(prognostic: np.ndarray, boundary: np.ndarray) -> int: |
| """ZeroGPU budget, measured live: ~0.3 s/step plus a small fixed cost.""" |
| return int(4 + 0.35 * boundary.shape[0]) |
|
|
|
|
| @spaces.GPU(duration=_gpu_duration) |
| def _rollout_gpu(prognostic: np.ndarray, boundary: np.ndarray) -> np.ndarray: |
| """Run the autoregressive rollout on GPU. |
| |
| Args: |
| prognostic: Normalized initial state, shape (1, 154, y, x). |
| boundary: Normalized forcing per model step, shape (n_steps, 8, y, x). |
| |
| Returns: |
| Normalized predictions, shape (n_steps, 154, y, x). |
| """ |
| x = torch.from_numpy(prognostic).to("cuda") |
| bnd = torch.from_numpy(boundary).to("cuda") |
| mask = LABEL_MASK.to("cuda") |
| outs = [] |
| with torch.no_grad(): |
| for k in range(bnd.shape[0]): |
| x = model.forward_once(x, bnd[k : k + 1], mask) |
| outs.append(x.float().cpu().numpy()) |
| return np.concatenate(outs, axis=0) |
|
|
|
|
| def _panel_frames(pred: np.ndarray, truth: np.ndarray, var_label: str, |
| start_index: int, path: str, fps: int) -> None: |
| """Write the prediction / truth / error animation to `path` as MP4.""" |
| _, cmap, units, short = VARIABLES[var_label] |
| finite = truth[np.isfinite(truth)] |
| vmin, vmax = np.percentile(finite, [1.0, 99.0]) |
| if short in ("SSH", "u 0 m", "v 0 m"): |
| lim = float(max(abs(vmin), abs(vmax))) |
| vmin, vmax = -lim, lim |
| diff = pred - truth |
| dlim = float(np.nanpercentile(np.abs(diff), 99.0)) or 1e-6 |
| extent = [store.lon[0], store.lon[-1], store.lat[0], store.lat[-1]] |
|
|
| field_cm = matplotlib.colormaps[cmap].copy() if isinstance(cmap, str) else cmap.copy() |
| diff_cm = matplotlib.colormaps["RdBu_r"].copy() |
| for c in (field_cm, diff_cm): |
| c.set_bad("#3a3f45") |
|
|
| fig, axes = plt.subplots(1, 3, figsize=(15.6, 4.0), dpi=110) |
| fig.patch.set_facecolor("#101418") |
| ims = [] |
| titles = [f"Samudra 2 — {short}", f"GFDL-OM4 (truth) — {short}", "Emulator − truth"] |
| for i, ax in enumerate(axes): |
| data = [pred, truth, diff][i][0] |
| kw = dict(cmap=field_cm, vmin=vmin, vmax=vmax) if i < 2 else dict( |
| cmap=diff_cm, vmin=-dlim, vmax=dlim) |
| im = ax.imshow(data, origin="lower", extent=extent, aspect="auto", |
| interpolation="nearest", **kw) |
| ims.append(im) |
| ax.set_title(titles[i], color="white", fontsize=10) |
| ax.tick_params(colors="#9aa4ad", labelsize=7) |
| for s in ax.spines.values(): |
| s.set_color("#3a3f45") |
| cb = fig.colorbar(im, ax=ax, fraction=0.032, pad=0.015) |
| cb.ax.tick_params(colors="#9aa4ad", labelsize=7) |
| cb.set_label(units, color="#9aa4ad", fontsize=8) |
| sup = fig.suptitle("", color="white", fontsize=11) |
| fig.tight_layout(rect=(0, 0, 1, 0.93)) |
|
|
| with imageio.get_writer(path, fps=fps, codec="libx264", quality=8, |
| macro_block_size=1, ffmpeg_log_level="error") as w: |
| for t in range(pred.shape[0]): |
| for im, data in zip(ims, (pred[t], truth[t], diff[t])): |
| im.set_data(data) |
| lead = (t + 1) * (STEP_DAYS // 2) |
| sup.set_text(f"valid {DATES[start_index + 2 + t]} · lead +{lead} days") |
| fig.canvas.draw() |
| frame = np.asarray(fig.canvas.buffer_rgba())[..., :3] |
| w.append_data(np.ascontiguousarray(frame)) |
| plt.close(fig) |
|
|
|
|
| def _skill_plot(pred: np.ndarray, truth: np.ndarray, persist: np.ndarray, |
| var_label: str, path: str) -> None: |
| """RMSE growth + global-mean drift, saved to `path` as PNG.""" |
| _, _, units, short = VARIABLES[var_label] |
| n = pred.shape[0] |
| lead = np.arange(1, n + 1) * (STEP_DAYS // 2) |
| ok = np.isfinite(truth[0]) & np.isfinite(pred[0]) |
| w = np.cos(np.deg2rad(store.lat))[:, None] * ok |
|
|
| def _rmse(a): |
| e = (a - truth) ** 2 |
| return np.sqrt(np.array([np.nansum(e[t] * w) / w.sum() for t in range(n)])) |
|
|
| def _mean(a): |
| return np.array([np.nansum(a[t] * w) / w.sum() for t in range(n)]) |
|
|
| fig, (ax1, ax2) = plt.subplots(1, 2, figsize=(11, 3.6), dpi=120) |
| fig.patch.set_facecolor("#101418") |
| ax1.plot(lead, _rmse(pred), "o-", color="#f59e0b", lw=2, ms=4, label="Samudra 2") |
| ax1.plot(lead, _rmse(persist), "s--", color="#94a3b8", lw=1.5, ms=3, |
| label="persistence") |
| ax1.set_ylabel(f"area-weighted RMSE [{units}]") |
| ax1.set_title(f"{short} — error growth", color="white", fontsize=11) |
| ax2.plot(lead, _mean(pred), "o-", color="#f59e0b", lw=2, ms=4, label="Samudra 2") |
| ax2.plot(lead, _mean(truth), "^-", color="#38bdf8", lw=2, ms=4, label="GFDL-OM4") |
| ax2.set_ylabel(f"global mean [{units}]") |
| ax2.set_title(f"{short} — global mean", color="white", fontsize=11) |
| for ax in (ax1, ax2): |
| ax.set_xlabel("lead time [days]") |
| ax.set_facecolor("#171c22") |
| ax.grid(alpha=0.18, color="#9aa4ad") |
| ax.tick_params(colors="#9aa4ad", labelsize=8) |
| ax.xaxis.label.set_color("#9aa4ad") |
| ax.yaxis.label.set_color("#9aa4ad") |
| for s in ax.spines.values(): |
| s.set_color("#3a3f45") |
| leg = ax.legend(fontsize=8, facecolor="#171c22", edgecolor="#3a3f45") |
| for txt in leg.get_texts(): |
| txt.set_color("#cbd5e1") |
| fig.tight_layout() |
| fig.savefig(path, facecolor=fig.get_facecolor()) |
| plt.close(fig) |
|
|
|
|
| def emulate( |
| start_date: str = DEFAULT_DATE, |
| n_steps: int = 12, |
| variable: str = DEFAULT_VARIABLE, |
| fps: int = 6, |
| progress=gr.Progress(), |
| ) -> tuple[str, str, str]: |
| """Roll Samudra 2 forward from a real GFDL-OM4 ocean state. |
| |
| Args: |
| start_date: Initial-condition date (5-daily, inside the held-out window). |
| n_steps: Number of model steps; each advances the ocean by 10 days. |
| variable: Ocean field to visualize. |
| fps: Frames per second of the output animation. |
| |
| Returns: |
| Path to the comparison MP4, path to the skill-score PNG, and a |
| markdown summary of the rollout. |
| """ |
| if variable not in VARIABLES: |
| raise gr.Error(f"Unknown variable: {variable}") |
| if start_date not in DATES: |
| raise gr.Error(f"{start_date} is not a valid OM4 timestamp.") |
| t0 = DATES.index(start_date) |
| n_steps = int(max(1, min(MAX_STEPS, n_steps))) |
| if t0 + 2 + 2 * n_steps > len(DATES): |
| n_steps = (len(DATES) - t0 - 2) // 2 |
| var, _, units, short = VARIABLES[variable] |
|
|
| progress(0.05, desc="Streaming initial condition from the OM4 store…") |
| prognostic = store.initial_prognostic(t0) |
| progress(0.20, desc="Streaming surface forcing…") |
| boundary = store.boundary_sequence(t0, n_steps) |
|
|
| progress(0.40, desc=f"Rolling out {n_steps} steps on GPU…") |
| t_gpu = time.time() |
| out = _rollout_gpu(prognostic, boundary) |
| gpu_s = time.time() - t_gpu |
|
|
| v = PROG_VARS.index(var) |
| pred = store.denormalize( |
| np.concatenate([out[k, [v, v + N_PROG]] for k in range(n_steps)]), var |
| ) |
|
|
| progress(0.70, desc="Fetching OM4 ground truth…") |
| truth = store.truth(var, t0, n_steps) |
| persist = np.broadcast_to(store.truth(var, t0 - HIST - 1, 1)[1:2], truth.shape) |
|
|
| progress(0.82, desc="Rendering animation…") |
| tmp = tempfile.mkdtemp() |
| video = os.path.join(tmp, "samudra2.mp4") |
| plot = os.path.join(tmp, "skill.png") |
| _panel_frames(pred, truth, variable, t0, video, int(fps)) |
| _skill_plot(pred, truth, persist, variable, plot) |
|
|
| ok = np.isfinite(pred) & np.isfinite(truth) |
| rmse = float(np.sqrt(np.mean((pred[ok] - truth[ok]) ** 2))) |
| pk = np.isfinite(persist) & np.isfinite(truth) |
| prmse = float(np.sqrt(np.mean((persist[pk] - truth[pk]) ** 2))) |
| end = DATES[t0 + 1 + 2 * n_steps] |
| summary = ( |
| f"**{short}** · initialised **{start_date}** → **{end}** " |
| f"({n_steps} model steps = {n_steps * STEP_DAYS} days, " |
| f"{2 * n_steps} frames)\n\n" |
| f"- RMSE vs GFDL-OM4: **{rmse:.4g} {units}** " |
| f"(persistence {prmse:.4g} {units} → " |
| f"**{100 * (1 - rmse / prmse):.0f}%** skill improvement)\n" |
| f"- Rollout on GPU: **{gpu_s:.1f} s** for {n_steps} steps " |
| f"({1000 * gpu_s / n_steps:.0f} ms/step)" |
| ) |
| progress(1.0, desc="Done") |
| return video, plot, summary |
|
|
|
|
| |
| CSS = """ |
| .dark .gradio-container { background: #0b0f14; } |
| #hdr h1 { margin-bottom: 0.15em; } |
| """ |
|
|
| with gr.Blocks(title="Samudra 2") as demo: |
| gr.Markdown( |
| """ |
| # 🌊 Samudra 2 — global ocean emulator |
| Roll [`M2LInES/Samudra2`](https://huggingface.co/M2LInES/Samudra2) forward from a |
| **real ocean state**. A ConvNeXt U-Net (84 M params) trained on GFDL-OM4 predicts |
| temperature, salinity, velocity and sea-surface height over 19 depth levels — |
| each model step advances the global ocean by **10 days**. |
| |
| The initial condition and forcing are streamed live from the public 1° OM4 store, |
| so every rollout is scored against the ocean model it emulates. |
| """, |
| elem_id="hdr", |
| ) |
| with gr.Row(): |
| with gr.Column(scale=1): |
| start_date = gr.Dropdown( |
| START_DATES, value=DEFAULT_DATE, label="Initial condition", |
| info="5-daily OM4 state, held-out period (2014-10-10 → 2022-12-24)", |
| ) |
| variable = gr.Dropdown( |
| list(VARIABLES), value=DEFAULT_VARIABLE, label="Field to visualize", |
| ) |
| n_steps = gr.Slider( |
| 1, MAX_STEPS, value=12, step=1, label="Rollout length (model steps)", |
| info="1 step = 10 days · 12 steps ≈ 4 months", |
| ) |
| run = gr.Button("Run rollout", variant="primary") |
| with gr.Accordion("Advanced", open=False): |
| fps = gr.Slider(2, 15, value=6, step=1, label="Animation frame rate") |
| with gr.Column(scale=2): |
| video = gr.Video(label="Emulator vs GFDL-OM4", autoplay=True, loop=True) |
| summary = gr.Markdown() |
| plot = gr.Image(label="Forecast skill", type="filepath") |
|
|
| inputs = [start_date, variable, n_steps, fps] |
|
|
| def _run(d: str = DEFAULT_DATE, v: str = DEFAULT_VARIABLE, n: int = 12, |
| f: int = 6, progress=gr.Progress()): |
| """Thin adapter so the UI order (date, variable, steps) drives `emulate`.""" |
| return emulate(d, n, v, f, progress) |
|
|
| run.click(_run, inputs=inputs, outputs=[video, plot, summary]) |
|
|
| gr.Examples( |
| examples=[ |
| ["2014-10-10", "Sea surface temperature (0 m)", 12], |
| ["2015-09-05", "Sea surface temperature (0 m)", 36], |
| ["2016-07-01", "Sea surface height", 18], |
| ["2018-03-24", "Sea surface salinity (0 m)", 12], |
| ["2019-05-18", "Temperature at 250 m", 24], |
| ["2017-01-28", "Zonal velocity at 0 m", 12], |
| ], |
| inputs=[start_date, variable, n_steps], |
| outputs=[video, plot, summary], |
| fn=_run, |
| cache_examples=True, |
| cache_mode="lazy", |
| ) |
|
|
| gr.Markdown( |
| "Model: [M2LInES/Samudra2](https://huggingface.co/M2LInES/Samudra2) (CC-BY-4.0) · " |
| "code [m2lines/Samudra](https://github.com/m2lines/Samudra) · " |
| "paper [arXiv:2606.02610](https://arxiv.org/abs/2606.02610) · " |
| "data: public 1° GFDL-OM4 store on the NYU OSN pod." |
| ) |
|
|
| demo.queue(max_size=12).launch(theme=gr.themes.Citrus(), css=CSS, mcp_server=True) |
|
|