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()