| |
| """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, |
| ) |
| |
| 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"] |
| ) |
| |
| BOUNDARY_VARS = ["tauuo", "tauvo", "hfds", "hfds_anomalies"] |
|
|
| N_PROG = len(PROG_VARS) |
| HIST = 1 |
| STEP_DAYS = 10 |
|
|
| _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() |
|
|
| |
| |
| |
| stats = np.load(HFDS_ANOM_STATS) |
| self._hfds_clim = stats["clim"].astype(np.float32) |
| 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 |
| |
| |
| 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)] |
| ) |
| 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) |
|
|
| |
| 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] |
|
|
| |
| 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) |
| 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) |
| anom = self._hfds_anomalies(raw[:, 2], t0, n_times)[:, None] |
| raw = np.concatenate([raw, anom], axis=1) |
|
|
| 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() |
|
|