File size: 7,241 Bytes
c21a022 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 | # SPDX-License-Identifier: Apache-2.0
"""Streaming access to the public 1° OM4 dataset used by Samudra 2.
The processed OM4 zarr stores live on the NYU OSN pod and are public-read:
https://nyu1.osn.mghpcc.org/m2lines-pubs/Samudra/v2026-07/om4_onedeg/
Only the handful of (time, y, x) chunks a rollout actually needs are pulled,
so a request moves tens of MB rather than the 92 GiB of the full store.
Channel layout, normalization and masking follow
`samudra.datasets.InferenceDataset` upstream:
* prognostic input = (hist+1=2 timesteps) x 77 variables = 154 channels
* boundary input = 2 timesteps x 4 variables = 8 channels
* model output = 154 channels = the next 2 timesteps of the 77 variables
so one model step advances the ocean state by 2 x 5 days = 10 days.
"""
from __future__ import annotations
import os
from concurrent.futures import ThreadPoolExecutor
from functools import lru_cache
import numpy as np
OSN_ENDPOINT = "https://nyu1.osn.mghpcc.org"
BUCKET_ROOT = "m2lines-pubs/Samudra/v2026-07/om4_onedeg"
LEVELS = 19
DEPTHS = (
2.5, 10.0, 22.5, 40.0, 65.0, 105.0, 165.0, 250.0, 375.0, 550.0,
775.0, 1050.0, 1400.0, 1850.0, 2400.0, 3100.0, 4000.0, 5000.0, 6000.0,
)
# `thermo_dynamic_all` prognostic variables, in upstream order.
PROG_VARS: list[str] = (
[f"uo_{i}" for i in range(LEVELS)]
+ [f"vo_{i}" for i in range(LEVELS)]
+ [f"thetao_{i}" for i in range(LEVELS)]
+ [f"so_{i}" for i in range(LEVELS)]
+ ["zos"]
)
# `tau_hfds_hfds_anom` boundary (forcing) variables.
BOUNDARY_VARS = ["tauuo", "tauvo", "hfds", "hfds_anomalies"]
N_PROG = len(PROG_VARS) # 77
HIST = 1 # samudra_om4_v2 default
STEP_DAYS = 10 # (hist + 1) * 5-day timesteps
_HERE = os.path.dirname(os.path.abspath(__file__))
HFDS_ANOM_STATS = os.path.join(_HERE, "hfds_anom_stats.npz")
def _level_of(var: str) -> int:
tail = var.rsplit("_", 1)[-1]
return int(tail) if tail.isdigit() else 0
class OM4Store:
"""Lazily-opened handle on the public 1° OM4 zarr store."""
def __init__(self, root: str = BUCKET_ROOT, endpoint: str = OSN_ENDPOINT):
import s3fs
import xarray as xr
fs = s3fs.S3FileSystem(anon=True, endpoint_url=endpoint)
self.ds = xr.open_zarr(s3fs.S3Map(root=f"{root}/OM4.zarr", s3=fs, check=False))
means = xr.open_zarr(
s3fs.S3Map(root=f"{root}/OM4_means.zarr", s3=fs, check=False)
).load()
stds = xr.open_zarr(
s3fs.S3Map(root=f"{root}/OM4_stds.zarr", s3=fs, check=False)
).load()
# hfds_anomalies is a derived channel (hfds minus its day-of-year
# climatology); the climatology + its normalization stats are shipped
# with the Space, precomputed from this very store.
stats = np.load(HFDS_ANOM_STATS)
self._hfds_clim = stats["clim"].astype(np.float32) # (73, y, x)
self._clim_doy = {int(d): i for i, d in enumerate(stats["dayofyear"])}
anom_mean, anom_std = float(stats["mean"]), float(stats["std"])
self.means = {v: float(means[v].values) for v in means.data_vars}
self.stds = {v: float(stds[v].values) for v in stds.data_vars}
self.means["hfds_anomalies"] = anom_mean
self.stds["hfds_anomalies"] = anom_std
self.time = self.ds.time.values
self.dayofyear = self.ds.time.dt.dayofyear.values
# `lat`/`lon` in the store are 2-D (y, x) curvilinear coords; the plot
# axes are the 1-D `y` / `x` cell centers.
self.lat = np.asarray(self.ds.y.values, np.float64)
self.lon = np.asarray(self.ds.x.values, np.float64)
self.masks = np.stack(
[self.ds[f"mask_{i}"].values.astype(bool) for i in range(LEVELS)]
) # (19, y, x)
self.shape = self.masks.shape[1:]
self.prog_mask = np.stack([self.masks[_level_of(v)] for v in PROG_VARS])
self.prog_means = np.array([self.means[v] for v in PROG_VARS], dtype=np.float32)
self.prog_stds = np.array([self.stds[v] for v in PROG_VARS], dtype=np.float32)
# ---------------------------------------------------------------- helpers
def date_str(self, index: int) -> str:
return str(self.time[index])[:10]
def _read(self, var: str, t0: int, n: int) -> np.ndarray:
"""(n, y, x) raw values for `var` over times [t0, t0+n)."""
return np.asarray(self.ds[var].isel(time=slice(t0, t0 + n)).values, np.float32)
def _read_many(self, variables: list[str], t0: int, n: int) -> np.ndarray:
"""(n, len(variables), y, x), fetched in parallel."""
with ThreadPoolExecutor(max_workers=16) as pool:
arrays = list(pool.map(lambda v: self._read(v, t0, n), variables))
return np.stack(arrays, axis=1)
def _hfds_anomalies(self, hfds: np.ndarray, t0: int, n: int) -> np.ndarray:
idx = [self._clim_doy[int(d)] for d in self.dayofyear[t0 : t0 + n]]
return hfds - self._hfds_clim[idx]
# ------------------------------------------------------------- public API
def initial_prognostic(self, t0: int) -> np.ndarray:
"""Normalized, masked (1, 154, y, x) initial state at times [t0, t0+1]."""
raw = self._read_many(PROG_VARS, t0, HIST + 1) # (2, 77, y, x)
norm = (raw - self.prog_means[None, :, None, None]) / self.prog_stds[
None, :, None, None
]
norm = np.nan_to_num(norm, nan=0.0)
norm = np.where(self.prog_mask[None], norm, 0.0)
return norm.reshape(1, (HIST + 1) * N_PROG, *self.shape).astype(np.float32)
def boundary_sequence(self, t0: int, n_steps: int) -> np.ndarray:
"""Normalized, masked (n_steps, 8, y, x) forcing for `n_steps` model steps."""
n_times = (HIST + 1) * n_steps
raw = self._read_many(["tauuo", "tauvo", "hfds"], t0, n_times) # (T, 3, y, x)
anom = self._hfds_anomalies(raw[:, 2], t0, n_times)[:, None]
raw = np.concatenate([raw, anom], axis=1) # (T, 4, y, x)
means = np.array([self.means[v] for v in BOUNDARY_VARS], np.float32)
stds = np.array([self.stds[v] for v in BOUNDARY_VARS], np.float32)
norm = (raw - means[None, :, None, None]) / stds[None, :, None, None]
norm = np.nan_to_num(norm, nan=0.0)
norm = np.where(self.masks[0][None, None], norm, 0.0)
return norm.reshape(n_steps, (HIST + 1) * len(BOUNDARY_VARS), *self.shape).astype(
np.float32
)
def truth(self, var: str, t0: int, n_steps: int) -> np.ndarray:
"""Raw (physical-unit) ground truth for `var` over the predicted times."""
n_times = (HIST + 1) * n_steps
raw = self._read(var, t0 + HIST + 1, n_times)
return np.where(self.masks[_level_of(var)][None], raw, np.nan)
def denormalize(self, channels: np.ndarray, var: str) -> np.ndarray:
"""Turn normalized model output for one variable into physical units."""
v = PROG_VARS.index(var)
out = channels * self.prog_stds[v] + self.prog_means[v]
return np.where(self.masks[_level_of(var)][None], out, np.nan)
@lru_cache(maxsize=1)
def get_store() -> OM4Store:
return OM4Store()
|