Skip to content
Draft
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
70 changes: 46 additions & 24 deletions climanet/dataset.py
Original file line number Diff line number Diff line change
Expand Up @@ -13,6 +13,16 @@
from typing import Tuple


def _as_numpy(array_like, dtype=None, copy=False):
"""Convert numpy/dask/xarray-backed arrays to numpy, computing lazily when needed."""
data = np.asarray(array_like)
if dtype is not None:
data = data.astype(dtype, copy=False)
if copy:
data = data.copy()
return data


class STDataset(Dataset):
"""Dataset for spatiotemporal patches.

Expand Down Expand Up @@ -60,7 +70,7 @@ def __init__(
self.spatial_dims = spatial_dims
self.patch_size = patch_size
self.input_da = input_da
self.monthly_da = monthly_da
self.monthly_da = monthly_da.compute() # Directly load monthly data
self.stride = stride if stride is not None else (patch_size[1], patch_size[2])

self.sh_embed_dim = sh_embed_dim
Expand Down Expand Up @@ -95,10 +105,8 @@ def __init__(
input_da, monthly_da, time_dim=time_dim
)

# Convert to tensor once — all __getitem__ calls use these
self.daily_t = torch.from_numpy(
daily_mt.values.astype(np.float32)
) # (M, T=31, H, W)
# Keep xarray objects (potentially dask-backed) lazy; materialize per patch in __getitem__
self.daily_mt = daily_mt # (M, T=31, H, W) or (M, T=31*24, H, W)
self.monthly_t = torch.from_numpy(
monthly_m.values.astype(np.float32)
) # (M, H, W)
Expand All @@ -121,13 +129,6 @@ def __init__(
else:
self.land_mask_t = None

# Precompute the NaN mask before filling NaNs
# daily_mask: True where NaN (i.e. missing ocean data, not land)
self.daily_nan_mask_t = torch.isnan(self.daily_t) # (M, T=31, H, W)

# NaNs will be filled with 0 in-place
self.daily_t.nan_to_num_(nan=0.0)

# Stats will be set later via set_stats() for train/test datasets
self.daily_mean = None
self.daily_std = None
Expand All @@ -136,8 +137,10 @@ def __init__(
_, ph, pw = self.patch_size
self._zero_land = torch.zeros(ph, pw, dtype=torch.bool)

# Precompute lazy index mapping for patches
M, H, W = self.daily_t.shape[0], self.daily_t.shape[2], self.daily_t.shape[3]
# Precompute lazy index mapping for patches from data shape metadata
M = int(self.daily_mt.sizes[self.daily_mt.dims[0]])
H = int(self.daily_mt.sizes[spatial_dims[0]])
W = int(self.daily_mt.sizes[spatial_dims[1]])
self.patch_indices = self._compute_patch_indices(M, H, W)

# Precompute geoposition and scale embeddings for patches
Expand Down Expand Up @@ -268,18 +271,25 @@ def __getitem__(self, idx):
m, i, j = self.patch_indices[idx]
pm, ph, pw = self.patch_size

# Extract spatial patch via slicing — faster than xarray indexing
# (M, T, H, W) -> (M,T,pH, pW)
daily_t_patch = self.daily_t[m : m + pm, :, i : i + ph, j : j + pw].unsqueeze(0)
# Compute only the requested patch from potentially dask-backed arrays.
daily_np_patch = self.daily_mt.isel(
{
self.daily_mt.dims[0]: slice(m, m + pm),
self.daily_mt.dims[2]: slice(i, i + ph),
self.daily_mt.dims[3]: slice(j, j + pw),
}
).values

# Build missing-data mask before NaN fill; True where NaN.
daily_nan_mask_t_patch = torch.from_numpy(np.isnan(daily_np_patch)).unsqueeze(0)

# Fill NaNs for model input.
np.nan_to_num(daily_np_patch, copy=False, nan=0.0)
daily_t_patch = torch.from_numpy(daily_np_patch).unsqueeze(0)

# (M, H, W) -> (M, pH, pW)
monthly_t_patch = self.monthly_t[m : m + pm, i : i + ph, j : j + pw]

# (M, T, H, W) -> (M, T, pH, pW)
daily_nan_mask_t_patch = self.daily_nan_mask_t[
m : m + pm, :, i : i + ph, j : j + pw
].unsqueeze(0)

if self.land_mask_t is not None:
land_t_patch = self.land_mask_t[i : i + ph, j : j + pw] # (H, W)
else:
Expand Down Expand Up @@ -331,14 +341,26 @@ def compute_stats(self, indices: list = None) -> Tuple[np.ndarray, np.ndarray]:
Tuple of (mean, std) arrays
"""
if indices is None:
data = self.monthly_t.numpy() # (M, H, W)
data = _as_numpy(
self.monthly_m.data, dtype=np.float32, copy=False
) # (M, H, W)
else:
# Stack selected spatial patches
pm, ph, pw = self.patch_size
patches = []
for idx in indices:
m, i, j = self.patch_indices[idx]
patch = self.monthly_t[m : m + pm, i : i + ph, j : j + pw].numpy()
patch = _as_numpy(
self.monthly_m.isel(
{
self.monthly_m.dims[0]: slice(m, m + pm),
self.monthly_m.dims[1]: slice(i, i + ph),
self.monthly_m.dims[2]: slice(j, j + pw),
}
).data,
dtype=np.float32,
copy=False,
)
patches.append(patch)
data = np.concatenate(patches, axis=-1)

Expand Down