diff --git a/src/samudra/aggregator/inference/main.py b/src/samudra/aggregator/inference/main.py index 57d8a144..e1aee182 100644 --- a/src/samudra/aggregator/inference/main.py +++ b/src/samudra/aggregator/inference/main.py @@ -7,8 +7,8 @@ import wandb import xarray as xr -from samudra.constants import TensorMap -from samudra.utils.data import Normalize, get_aggregator_dicts +from samudra.constants import DataLayout +from samudra.utils.data import BatchPreprocessor, get_aggregator_dicts from samudra.utils.output import ModelInferenceOutput from samudra.utils.wandb import Metrics, MetricsDict @@ -30,8 +30,8 @@ def __init__( area_weights: torch.Tensor, wet: torch.Tensor, num_prognostic_channels: int, - normalize: Normalize, - tensor_map: TensorMap, + preprocessor: BatchPreprocessor, + data_layout: DataLayout, record_step_20: bool = True, log_global_mean_time_series: bool = True, log_global_mean_norm_time_series: bool = True, @@ -47,8 +47,8 @@ def __init__( area_weights: Area weights for the data. wet: Wet mask for the data. num_prognostic_channels: Number of prognostic channels in the data. - normalize: Normalization helper for prognostic channels. - tensor_map: Mapping from prognostic variables to tensor channels. + preprocessor: Normalization helper for prognostic channels. + data_layout: Mapping from prognostic variables to tensor channels. record_step_20: Whether to record the mean of the 20th steps. log_global_mean_time_series: Whether to log global mean time series metrics. log_global_mean_norm_time_series: Whether to log the normalized global mean @@ -103,8 +103,8 @@ def __init__( if name not in ["mean", "mean_norm"] } self._n_timesteps_seen = 0 - self._normalize = normalize - self._tensor_map = tensor_map + self._preprocessor = preprocessor + self._data_layout = data_layout self.num_prognostic_channels = num_prognostic_channels self.hist = hist self.wet = wet @@ -123,8 +123,8 @@ def record_batch(self, data: ModelInferenceOutput): assert data.prediction.shape[0] == total_len // (self.hist + 1) target_norm_dict, target_unnorm_dict = get_aggregator_dicts( data.target, - normalize=self._normalize, - tensor_map=self._tensor_map, + preprocessor=self._preprocessor, + data_layout=self._data_layout, wet=self.wet, long_rollout=True, input_type="prognostic", @@ -133,8 +133,8 @@ def record_batch(self, data: ModelInferenceOutput): ) gen_norm_dict, gen_unnorm_dict = get_aggregator_dicts( data.prediction, - normalize=self._normalize, - tensor_map=self._tensor_map, + preprocessor=self._preprocessor, + data_layout=self._data_layout, wet=self.wet, long_rollout=True, input_type="prognostic", @@ -177,8 +177,8 @@ def record_initial_prognostic( data_norm_dict, data_unnorm_dict = get_aggregator_dicts( initial_prognostic, - normalize=self._normalize, - tensor_map=self._tensor_map, + preprocessor=self._preprocessor, + data_layout=self._data_layout, wet=self.wet, long_rollout=True, input_type="input", diff --git a/src/samudra/aggregator/loss.py b/src/samudra/aggregator/loss.py index 6bca1e51..5b488c27 100644 --- a/src/samudra/aggregator/loss.py +++ b/src/samudra/aggregator/loss.py @@ -5,19 +5,19 @@ import torch -from samudra.constants import TensorMap +from samudra.constants import DataLayout def get_depth_loss_dict( label: str, loss_per_channel: torch.Tensor, *, - tensor_map: TensorMap, + data_layout: DataLayout, ) -> dict[str, torch.Tensor]: metrics = {} - for depth in tensor_map.DEPTH_SET: + for depth in data_layout.depths: metrics[f"{label}/loss/depth/depth_{depth}_loss"] = loss_per_channel[ - tensor_map.DP_3D_IDX[depth] + data_layout.depth_indices[depth] ].mean() return metrics @@ -26,12 +26,12 @@ def get_variable_loss_dict( label: str, loss_per_channel: torch.Tensor, *, - tensor_map: TensorMap, + data_layout: DataLayout, ) -> dict[str, torch.Tensor]: metrics = {} - for variable in tensor_map.VAR_SET: + for variable in data_layout.variables: metrics[f"{label}/loss/variable/{variable}_loss"] = loss_per_channel[ - tensor_map.VAR_3D_IDX[variable] + data_layout.variable_indices[variable] ].mean() return metrics @@ -40,23 +40,23 @@ def get_channel_loss_dict( label: str, loss_per_channel: torch.Tensor, *, - tensor_map: TensorMap, + data_layout: DataLayout, loss_name: str = "loss", ) -> dict[str, torch.Tensor]: - return get_channel_dict(label, loss_name, loss_per_channel, tensor_map=tensor_map) + return get_channel_dict(label, loss_name, loss_per_channel, data_layout=data_layout) def get_channel_loss_scale_dict( label: str, loss_scale_per_channel: torch.Tensor, *, - tensor_map: TensorMap, + data_layout: DataLayout, ) -> dict[str, torch.Tensor]: return get_channel_dict( label, "loss_scale", loss_scale_per_channel, - tensor_map=tensor_map, + data_layout=data_layout, ) @@ -65,9 +65,9 @@ def get_channel_dict( measure: str, per_channel: torch.Tensor, *, - tensor_map: TensorMap, + data_layout: DataLayout, ) -> dict[str, torch.Tensor]: metrics = {} - for i, channel in enumerate(tensor_map.prognostic_var_names): + for i, channel in enumerate(data_layout.prognostic_var_names): metrics[f"{prefix}/{measure}/channel/{channel}_{measure}"] = per_channel[i] return metrics diff --git a/src/samudra/aggregator/main.py b/src/samudra/aggregator/main.py index b09cb465..b27040fe 100644 --- a/src/samudra/aggregator/main.py +++ b/src/samudra/aggregator/main.py @@ -18,14 +18,14 @@ from samudra.aggregator.validate.reduced import MeanAggregator from samudra.aggregator.validate.snapshot import SnapshotAggregator from samudra.aggregator.validate.sub_aggregator import ValidateSubAggregator -from samudra.constants import TensorMap -from samudra.utils.data import Normalize +from samudra.constants import DataLayout +from samudra.utils.data import BatchPreprocessor class Aggregator: @staticmethod - def get_train_aggregator(tensor_map: TensorMap) -> TrainAggregator: - return TrainAggregator(tensor_map) + def get_train_aggregator(data_layout: DataLayout) -> TrainAggregator: + return TrainAggregator(data_layout) @staticmethod def get_validation_aggregator( @@ -33,8 +33,8 @@ def get_validation_aggregator( hist: int, area_weights: torch.Tensor, num_prognostic_channels: int, - tensor_map: TensorMap, - normalize: Normalize, + data_layout: DataLayout, + preprocessor: BatchPreprocessor, *, include_image_aggregators: bool = True, ) -> ValidateAggregator: @@ -53,8 +53,8 @@ def get_validation_aggregator( val_aggregators, hist=hist, num_prognostic_channels=num_prognostic_channels, - tensor_map=tensor_map, - normalize=normalize, + data_layout=data_layout, + preprocessor=preprocessor, ) @staticmethod @@ -65,8 +65,8 @@ def get_inline_inference_aggregator( area_weights: torch.Tensor, wet: torch.Tensor, num_prognostic_channels: int, - tensor_map: TensorMap, - normalize: Normalize, + data_layout: DataLayout, + preprocessor: BatchPreprocessor, channel_mean_names: list[str] | None = None, ) -> InferenceEvaluatorAggregator: return InferenceEvaluatorAggregator( @@ -76,8 +76,8 @@ def get_inline_inference_aggregator( area_weights=area_weights, wet=wet, num_prognostic_channels=num_prognostic_channels, - normalize=normalize, - tensor_map=tensor_map, + preprocessor=preprocessor, + data_layout=data_layout, record_step_20=(n_timesteps > 20), log_global_mean_time_series=False, log_global_mean_norm_time_series=False, @@ -92,8 +92,8 @@ def get_standalone_inference_aggregator( area_weights: torch.Tensor, wet: torch.Tensor, num_prognostic_channels: int, - tensor_map: TensorMap, - normalize: Normalize, + data_layout: DataLayout, + preprocessor: BatchPreprocessor, channel_mean_names: list[str] | None = None, ) -> InferenceEvaluatorAggregator: return InferenceEvaluatorAggregator( @@ -103,8 +103,8 @@ def get_standalone_inference_aggregator( area_weights=area_weights, wet=wet, num_prognostic_channels=num_prognostic_channels, - normalize=normalize, - tensor_map=tensor_map, + preprocessor=preprocessor, + data_layout=data_layout, record_step_20=(n_timesteps > 20), log_global_mean_time_series=True, log_global_mean_norm_time_series=True, diff --git a/src/samudra/aggregator/train.py b/src/samudra/aggregator/train.py index 70c894fa..8df720d6 100644 --- a/src/samudra/aggregator/train.py +++ b/src/samudra/aggregator/train.py @@ -10,7 +10,7 @@ get_depth_loss_dict, get_variable_loss_dict, ) -from samudra.constants import TensorMap +from samudra.constants import DataLayout from samudra.utils.distributed import all_reduce_mean from samudra.utils.output import TrainBatchOutput from samudra.utils.wandb import Metrics @@ -19,8 +19,8 @@ class TrainAggregator: """Aggregates train statistics for an epoch.""" - def __init__(self, tensor_map: TensorMap): - self.tensor_map = tensor_map + def __init__(self, data_layout: DataLayout): + self.data_layout = data_layout self._n_batches = 0 self._loss = torch.tensor(torch.nan) self._loss_per_channel = torch.tensor(torch.nan) @@ -41,13 +41,13 @@ def get_logs(self, label: str = "train") -> Metrics: loss_per_channel = self._loss_per_channel / self._n_batches depth_loss_dict = get_depth_loss_dict( - label, loss_per_channel, tensor_map=self.tensor_map + label, loss_per_channel, data_layout=self.data_layout ) var_loss_dict = get_variable_loss_dict( - label, loss_per_channel, tensor_map=self.tensor_map + label, loss_per_channel, data_layout=self.data_layout ) channel_loss_dict = get_channel_loss_dict( - label, loss_per_channel, tensor_map=self.tensor_map + label, loss_per_channel, data_layout=self.data_layout ) logs = { f"{label}/mean/loss": loss, diff --git a/src/samudra/aggregator/validate/main.py b/src/samudra/aggregator/validate/main.py index 0f920db6..ec96c164 100644 --- a/src/samudra/aggregator/validate/main.py +++ b/src/samudra/aggregator/validate/main.py @@ -7,8 +7,8 @@ from samudra.aggregator.train import TrainAggregator from samudra.aggregator.validate.sub_aggregator import ValidateSubAggregator -from samudra.constants import TensorMap -from samudra.utils.data import Normalize, get_aggregator_dicts +from samudra.constants import DataLayout +from samudra.utils.data import BatchPreprocessor, get_aggregator_dicts from samudra.utils.output import ValBatchOutput from samudra.utils.wandb import Metrics, MetricsDict @@ -22,14 +22,14 @@ def __init__( hist: int, num_prognostic_channels: int, *, - tensor_map: TensorMap, - normalize: Normalize, + data_layout: DataLayout, + preprocessor: BatchPreprocessor, ): - super().__init__(tensor_map) + super().__init__(data_layout) self._aggregators = aggregators self.hist = hist self.num_prognostic_channels = num_prognostic_channels - self.normalize = normalize + self.preprocessor = preprocessor # TODO(jder): we could remove this by moving from inheritance # to composition with the TrainAggregator functionality. @@ -46,7 +46,7 @@ def record_validation_batch(self, batch: ValBatchOutput): if not self._aggregators: return - # Translate the GridContext mask by removing history. + # Translate the BatchGrid mask by removing history. target_data = batch.target_data # [B, C*(hist+1), H, W] wet = batch.ctx.label_mask # [C*(hist+1), H, W] assert wet.shape == target_data.shape[1:], ( @@ -66,8 +66,8 @@ def record_validation_batch(self, batch: ValBatchOutput): assert target_data.shape[1] == self.num_prognostic_channels target_data_dict, target_data_unnorm_dict = get_aggregator_dicts( target_data, - normalize=self.normalize, - tensor_map=self.tensor_map, + preprocessor=self.preprocessor, + data_layout=self.data_layout, wet=wet, long_rollout=False, input_type="prognostic", @@ -77,8 +77,8 @@ def record_validation_batch(self, batch: ValBatchOutput): gen_data_dict, gen_data_unnorm_dict = get_aggregator_dicts( batch.gen_data, - normalize=self.normalize, - tensor_map=self.tensor_map, + preprocessor=self.preprocessor, + data_layout=self.data_layout, wet=wet, long_rollout=False, input_type="prognostic", @@ -87,8 +87,8 @@ def record_validation_batch(self, batch: ValBatchOutput): ) input_data_dict, input_data_unnorm_dict = get_aggregator_dicts( batch.input_data, - normalize=self.normalize, - tensor_map=self.tensor_map, + preprocessor=self.preprocessor, + data_layout=self.data_layout, wet=wet, long_rollout=False, input_type="input", diff --git a/src/samudra/config.py b/src/samudra/config.py index 58c8df82..a781325a 100644 --- a/src/samudra/config.py +++ b/src/samudra/config.py @@ -20,11 +20,11 @@ from samudra.config_base import BaseConfig, TopLevelConfig from samudra.constants import ( - DatasetSpec, + DataLayout, + GridSize, GridType, LoaderVersion, - build_llc_spec, - build_om4_spec, + build_om4_layout, ) from samudra.models import Samudra, SamudraMini, SamudraMulti from samudra.models.base import BaseModel @@ -50,7 +50,16 @@ ) from samudra.models.modules.blocks import ZonallyPeriodicBilinearUpsample from samudra.models.modules.encoder import patch_from -from samudra.utils.data import DataContainer, DataSource, DataSourceSplits +from samudra.utils.data import ( + CanonicalSource, + DataBundle, + SourceSplits, + compute_anomalies, + flatten_masks, + get_anomalies_vars, + with_lat_lon_coords, + with_level_index_vars, +) from samudra.utils.llc import canonicalize_llc_datasets from samudra.utils.location import LocalLocation, Location, ResolvedLocation from samudra.utils.loss import ( @@ -186,18 +195,13 @@ class BaseDataSourceConfig[SourceTimeConfigT: TimeConfig](BaseConfig, abc.ABC): description="Location of the data standard deviations; " + LOCATION_DOCS ) - @property - @abc.abstractmethod - def dataset_spec(self) -> DatasetSpec: - raise NotImplementedError - @abc.abstractmethod def canonicalize_datasets( self, data: xr.Dataset, means: xr.Dataset, stds: xr.Dataset, - ) -> tuple[xr.Dataset, xr.Dataset, xr.Dataset]: + ) -> tuple[xr.Dataset, xr.Dataset, xr.Dataset, DataLayout]: raise NotImplementedError def build( @@ -206,7 +210,7 @@ def build( *, use_dask: bool, is_primary: bool, - ) -> DataSourceSplits: + ) -> SourceSplits: source = self._build_source( data_root, turn_on_dask=use_dask, @@ -224,11 +228,11 @@ def build( assert len(self.inference_times) == 1, ( "multiple inference time ranges have been deprecated" ) - inference_source = full_inference_source.slice(self.inference_times[0]) + inference_source = full_inference_source.slice_time(self.inference_times[0]) - return DataSourceSplits( - train=source.slice(self.train_time), - val=source.slice(self.val_time), + return SourceSplits( + train=source.slice_time(self.train_time), + val=source.slice_time(self.val_time), inference=inference_source, ) @@ -237,7 +241,7 @@ def _build_source( data_root: ResolvedLocation, *, turn_on_dask: bool, - ) -> DataSource: + ) -> CanonicalSource: resolved_data_location = data_root.resolve(self.data_location) resolved_means_location = data_root.resolve(self.data_means_location) resolved_stds_location = data_root.resolve(self.data_stds_location) @@ -246,20 +250,18 @@ def _build_source( data = resolved_data_location.open(chunks) means = resolved_means_location.open(chunks) stds = resolved_stds_location.open(chunks) - data, means, stds = self.canonicalize_datasets( + data, means, stds, data_layout = self.canonicalize_datasets( data, means, stds, ) - dataset_spec = self.dataset_spec - - source = DataSource.from_datasets( + source = CanonicalSource.from_datasets( data, means, stds, - dataset_spec=dataset_spec, - prognostic_var_names=dataset_spec.prognostic_var_names, - boundary_var_names=dataset_spec.boundary_var_names, + data_layout=data_layout, + prognostic_var_names=data_layout.prognostic_var_names, + boundary_var_names=data_layout.boundary_var_names, name=f"{resolved_data_location}-{turn_on_dask}", ) return source @@ -311,14 +313,6 @@ class Om4DataSourceConfig(BaseDataSourceConfig[Om4TimeConfig]): boundary_vars_key: str = "tau_hfds" grid_type: GridType = "gaussian" - @property - def dataset_spec(self) -> DatasetSpec: - return build_om4_spec( - self.prognostic_vars_key, - self.boundary_vars_key, - grid_type=self.grid_type, - ) - @pydantic.model_validator(mode="after") def validate_time_splits(self) -> Self: if self.train_time.overlaps(self.val_time): @@ -333,9 +327,44 @@ def canonicalize_datasets( data: xr.Dataset, means: xr.Dataset, stds: xr.Dataset, - ) -> tuple[xr.Dataset, xr.Dataset, xr.Dataset]: - # TODO: move OM4 canonicalization in here instead of in validation afte. - return data, means, stds + ) -> tuple[xr.Dataset, xr.Dataset, xr.Dataset, DataLayout]: + """Convert raw flat or compact OM4 xarray inputs to canonical channels.""" + data_layout = build_om4_layout( + self.prognostic_vars_key, + self.boundary_vars_key, + grid_type=self.grid_type, + ) + data = data.copy() + means = means.copy() + stds = stds.copy() + + data = with_lat_lon_coords(data) + data = with_level_index_vars(data, depth_levels=data_layout.depth_levels) + means = with_level_index_vars(means, depth_levels=data_layout.depth_levels) + stds = with_level_index_vars(stds, depth_levels=data_layout.depth_levels) + data = flatten_masks(data) + + anomalies_vars = get_anomalies_vars(data_layout.boundary_var_names) + if anomalies_vars: + data, means, stds = compute_anomalies(data, means, stds, anomalies_vars) + + def expand_levels(dataset: xr.Dataset) -> xr.Dataset: + canonical = xr.Dataset(attrs=dataset.attrs) + for coord in ("time", "lat", "lon"): + if coord in dataset.coords: + canonical = canonical.assign_coords({coord: dataset.coords[coord]}) + for name, variable in dataset.data_vars.items(): + if "lev" not in variable.dims: + canonical[str(name)] = variable + continue + for level in range(variable.sizes["lev"]): + canonical[f"{name}_{level}"] = variable.isel(lev=level, drop=True) + return canonical + + canonical_data = expand_levels(data) + canonical_means = expand_levels(means) + canonical_stds = expand_levels(stds) + return canonical_data, canonical_means, canonical_stds, data_layout class LlcDataSourceConfig(BaseDataSourceConfig[LlcTimeConfig]): @@ -348,13 +377,6 @@ class LlcDataSourceConfig(BaseDataSourceConfig[LlcTimeConfig]): j_start: int = Field(default=0, ge=0) j_end: int = Field(default=720, gt=0) - @property - def dataset_spec(self) -> DatasetSpec: - return build_llc_spec( - self.prognostic_vars_key, - self.boundary_vars_key, - ) - @pydantic.model_validator(mode="after") def validate_time_splits(self) -> Self: if self.train_time.overlaps(self.val_time): @@ -377,7 +399,7 @@ def canonicalize_datasets( data: xr.Dataset, means: xr.Dataset, stds: xr.Dataset, - ) -> tuple[xr.Dataset, xr.Dataset, xr.Dataset]: + ) -> tuple[xr.Dataset, xr.Dataset, xr.Dataset, DataLayout]: return canonicalize_llc_datasets( data, means, @@ -387,7 +409,8 @@ def canonicalize_datasets( i_end=self.i_end, j_start=self.j_start, j_end=self.j_end, - dataset_spec=self.dataset_spec, + prognostic_vars_key=self.prognostic_vars_key, + boundary_vars_key=self.boundary_vars_key, ) @@ -415,13 +438,9 @@ class DataConfig(BaseConfig): def build( self, data_root: ResolvedLocation, - ) -> DataContainer: + ) -> DataBundle: loader_version = LoaderVersion(self.loader_version) use_dask = loader_version != LoaderVersion.OM4_TORCH - dataset_spec = self.sources[0].dataset_spec - assert all(source.dataset_spec == dataset_spec for source in self.sources), ( - "All data sources must use the same dataset spec" - ) source_splits = [ source_cfg.build( @@ -433,13 +452,17 @@ def build( ] train_sources = [splits.train for splits in source_splits] val_sources = [splits.val for splits in source_splits] + primary_source = train_sources[0] + data_layout = primary_source.data_layout + if any(source.data_layout != data_layout for source in train_sources[1:]): + raise ValueError("All data sources must use the same data layout") - return DataContainer( + return DataBundle( train_sources=train_sources, val_sources=val_sources, inference_source=source_splits[0].inference, loader_version=loader_version, - dataset_spec=dataset_spec, + data_layout=data_layout, ) @@ -825,7 +848,7 @@ def build( boundary_channels: int, out_channels: int, hist: int, - srcs: list[DataSource], + grid_sizes: list[GridSize], ) -> BaseModel: pass @@ -847,13 +870,12 @@ def build( boundary_channels: int, out_channels: int, hist: int, - srcs: list[DataSource], + grid_sizes: list[GridSize], ) -> Samudra: - if len(srcs) != 1: + if len(grid_sizes) != 1: raise ValueError( "Samudra only supports training at a single scale! Please configure exactly one data source." ) - src = srcs[0] in_channels = prog_channels + boundary_channels total_in_channels = ( in_channels + self.pos_channels + (3 if self.add_3d_coordinates else 0) @@ -873,7 +895,7 @@ def build( pos_channels=self.pos_channels, add_3d_coordinates=add_3d_coordinates, hist=hist, - grid_size=src.grid_size, + grid_size=grid_sizes[0], gradient_detach_interval=self.gradient_detach_interval, use_bfloat16=self.use_bfloat16, ) @@ -905,15 +927,14 @@ def build( boundary_channels: int, out_channels: int, hist: int, - srcs: list[DataSource], + grid_sizes: list[GridSize], ) -> SamudraMulti: assert len(self.patch_extent) == 2, "patch_extent must be a pair of floats." extent = self.patch_extent[0], self.patch_extent[1] - all_grid_sizes = [s.grid_size for s in srcs] max_lat_size, max_lon_size = ( - max(g[0] for g in all_grid_sizes), - max(g[1] for g in all_grid_sizes), + max(g[0] for g in grid_sizes), + max(g[1] for g in grid_sizes), ) impl = self.perceiver_implementation @@ -999,7 +1020,7 @@ def build( boundary_channels: int, out_channels: int, hist: int, - srcs: list[DataSource], + grid_sizes: list[GridSize], ) -> SamudraMini: if self.add_3d_coordinates: raise ValueError( diff --git a/src/samudra/constants.py b/src/samudra/constants.py index a3e33eab..76f0b937 100644 --- a/src/samudra/constants.py +++ b/src/samudra/constants.py @@ -5,7 +5,7 @@ import dataclasses import enum import logging -from typing import Literal, Self +from typing import Literal, NamedTuple, Self logger = logging.getLogger(__name__) @@ -32,9 +32,12 @@ # So, we'll leave this default and use symbolic axes locally. type Input = Float[Grid, "*batch total_vars"] -Example = tuple[ - Prognostic, Boundary, Prognostic -] # (prognostic_input, boundary_input, label) + +class RolloutStep(NamedTuple): + prognostic: Prognostic + boundary: Boundary + label: Prognostic + GridMask = Bool[Tensor, "lat lon"] PrognosticMask = Bool[GridMask, "prognostic_vars"] @@ -51,41 +54,96 @@ MAX_TRAIN_MODEL_STEPS_FORWARD = 200 -DatasetType = Literal["om4", "llc"] - # Horizontal grid geometry. "gaussian" is a regular/rectilinear lat-lon grid whose # 2D lat/lon are the outer product of the 1D axes, so they can be reconstructed by # broadcasting. "tripolar" is curvilinear: lat/lon vary along both horizontal dims # and cannot be rebuilt by broadcasting, so the real 2D coordinates must be carried # with the data. Downstream code that reconstructs geometry must branch on this. GridType = Literal["gaussian", "tripolar"] - PrognosticVarNames = list[str] BoundaryVarNames = list[str] @dataclasses.dataclass(frozen=True) -class DatasetSpec: - type: DatasetType +class DataLayout: + """Canonical channel layout, physical metadata, and tensor index mappings. + + Raw dataset conventions belong to source-specific canonicalizers, not here. + """ + depth_levels: tuple[float, ...] depth_thickness: tuple[float, ...] - mask_vars: tuple[str, ...] - mask_all_levels_var: str - seconds_per_time_step: int prognostic_var_names: PrognosticVarNames boundary_var_names: BoundaryVarNames default_metadata: dict[str, dict[str, str]] ocean_heat_temperature_var: str - surface_heat_flux_var: str grid_type: GridType = "gaussian" + variable_indices: dict[str, torch.Tensor] = dataclasses.field( + init=False, repr=False, compare=False + ) + depth_indices: dict[str, torch.Tensor] = dataclasses.field( + init=False, repr=False, compare=False + ) + variables: list[str] = dataclasses.field(init=False, compare=False) + depths: list[str] = dataclasses.field(init=False, compare=False) + dz: torch.Tensor = dataclasses.field(init=False, repr=False, compare=False) def __post_init__(self) -> None: if len(self.depth_levels) != len(self.depth_thickness): raise ValueError( "depth_levels and depth_thickness must have the same length" ) - if len(self.depth_levels) != len(self.mask_vars): - raise ValueError("depth_levels and mask_vars must have the same length") + + def split_channel(channel: str) -> tuple[str, str | None]: + name, separator, suffix = channel.rpartition("_") + return (name, suffix) if separator and suffix.isdigit() else (channel, None) + + split_channels = [ + split_channel(channel) for channel in self.prognostic_var_names + ] + var_set_2d = [name for name, depth in split_channels if depth is None] + var_set = list(dict.fromkeys(name for name, _ in split_channels)) + depth_set = list(self.depth_i_levels[: self.num_prognostic_depth_levels]) + + var_indices = { + name: torch.tensor( + [ + index + for index, (channel_name, _) in enumerate(split_channels) + if name == channel_name + ], + dtype=torch.int32, + ) + for name in var_set + } + depth_indices = { + depth: torch.tensor( + [ + index + for index, (_, channel_depth) in enumerate(split_channels) + if channel_depth == depth + ], + dtype=torch.int32, + ) + for depth in depth_set + } + if depth_set and var_set_2d: + depth_indices[depth_set[0]] = torch.cat( + [ + depth_indices[depth_set[0]], + *(var_indices[name] for name in var_set_2d), + ] + ) + + object.__setattr__(self, "variable_indices", var_indices) + object.__setattr__(self, "depth_indices", depth_indices) + object.__setattr__(self, "variables", var_set) + object.__setattr__(self, "depths", depth_set) + object.__setattr__( + self, + "dz", + torch.tensor(self.depth_thickness[: self.num_prognostic_depth_levels]), + ) @property def depth_i_levels(self) -> tuple[str, ...]: @@ -100,6 +158,16 @@ def num_prognostic_depth_levels(self) -> int: depth_indices.append(int(suffix)) return max(depth_indices) + 1 if depth_indices else 1 + def to(self, device: torch.device) -> Self: + """Return a copy whose tensor indices live on ``device``.""" + moved = dataclasses.replace(self) + for mapping_name in ("variable_indices", "depth_indices"): + mapping = getattr(moved, mapping_name) + for key, value in mapping.items(): + mapping[key] = value.to(device) + object.__setattr__(moved, "dz", moved.dz.to(device)) + return moved + # Experiment prognostic and boundary variables # Assumption that all 3D variables are appended with depth_i_levels @@ -112,7 +180,7 @@ def num_prognostic_depth_levels(self) -> int: def _select_var_names( vars_by_key: dict[str, list[str]], key: str, - dataset_type: DatasetType, + dataset_type: str, var_kind: str, ) -> list[str]: try: @@ -125,13 +193,12 @@ def _select_var_names( ) from exc -def build_om4_spec( +def build_om4_layout( prognostic_vars_key: str = "thermo_dynamic_all", boundary_vars_key: str = "tau_hfds", grid_type: GridType = "gaussian", -) -> DatasetSpec: - return DatasetSpec( - type="om4", +) -> DataLayout: + return DataLayout( depth_levels=( 2.5, 10.0, @@ -174,9 +241,6 @@ def build_om4_spec( 1000.0, 1000.0, ), - mask_vars=tuple(f"mask_{i}" for i in range(19)), - mask_all_levels_var="wetmask", - seconds_per_time_step=5 * 24 * 60 * 60, prognostic_var_names=_select_var_names( { "thetao_1": ["thetao_0"], @@ -255,17 +319,15 @@ def build_om4_spec( }, }, ocean_heat_temperature_var="thetao", - surface_heat_flux_var="hfds", grid_type=grid_type, ) -def build_llc_spec( +def build_llc_layout( prognostic_vars_key: str = "single_1", boundary_vars_key: str = "single_1", -) -> DatasetSpec: - return DatasetSpec( - type="llc", +) -> DataLayout: + return DataLayout( depth_levels=( 0.5, 1.57, @@ -372,9 +434,6 @@ def build_llc_spec( 45.46, 54.405, ), - mask_vars=tuple(f"wetmask_{i}" for i in range(51)), - mask_all_levels_var="wetmask", - seconds_per_time_step=60, prognostic_var_names=_select_var_names( { "single_1": ["Theta_0"], @@ -434,7 +493,6 @@ def build_llc_spec( }, }, ocean_heat_temperature_var="Theta", - surface_heat_flux_var="oceQnet", # LLC (lat-lon-cap) is curvilinear, so its 2D geometry can't be broadcast # from 1D axes -- same broadcast-unsafe class as the tripolar grid. grid_type="tripolar", @@ -443,7 +501,7 @@ def build_llc_spec( def construct_metadata( data: xr.Dataset, - dataset_spec: DatasetSpec, + data_layout: DataLayout, ) -> dict[str, dict[str, str]]: metadata = {} for var in data.variables: @@ -453,12 +511,10 @@ def construct_metadata( "units": data[var].units, } except AttributeError: - if var in dataset_spec.default_metadata.keys(): - metadata[str(var)] = dataset_spec.default_metadata[str(var)] - elif ( - key := str(var).split("_")[0] - ) in dataset_spec.default_metadata.keys(): - metadata[str(var)] = dataset_spec.default_metadata[key] + if var in data_layout.default_metadata.keys(): + metadata[str(var)] = data_layout.default_metadata[str(var)] + elif (key := str(var).split("_")[0]) in data_layout.default_metadata.keys(): + metadata[str(var)] = data_layout.default_metadata[key] else: logger.info(f"{var} does not have any default metadata") metadata[str(var)] = { @@ -471,100 +527,3 @@ def construct_metadata( class LoaderVersion(enum.Enum): OM4_TORCH = "om4-torch" - - -class TensorMap: - def __init__( - self, - dataset_spec: DatasetSpec, - ): - """ - Maps input variables / depth levels to their indices in the input tensor. - - VAR_3D_IDX maps the input variables to their indices in the input tensor - DP_3D_IDX maps the depth levels to their indices in the input tensor - """ - self.dataset_spec = dataset_spec - self.VAR_3D_IDX: dict[str, torch.Tensor] = {} - self.DP_3D_IDX: dict[str, torch.Tensor] = {} - - self.INPT_BOUNDARY_IDX: dict[str, torch.Tensor] = {} - self.VAR_SET_2D = [] - self.VAR_SET_3D = [] - self.prognostic_var_names = dataset_spec.prognostic_var_names - self.boundary_var_names = dataset_spec.boundary_var_names - for out in self.prognostic_var_names: - var_split = out.split("_") - if len(var_split) == 1: - self.VAR_SET_2D.append(var_split[0]) - else: - self.VAR_SET_3D.append(var_split[0]) - - # Consistent order of variables - self.VAR_SET = list( - dict.fromkeys([out.split("_")[0] for out in self.prognostic_var_names]) - ) - - levels = dataset_spec.num_prognostic_depth_levels - - self.DEPTH_SET = list(dataset_spec.depth_i_levels[:levels]) - self.dz = torch.tensor(dataset_spec.depth_thickness[:levels]) - - self._populate_var_3d_idx() - self._populate_dp_3d_idx() - self._populate_boundary_idx() - - def _populate_var_3d_idx(self): - for kt in self.VAR_SET: - self.VAR_3D_IDX[kt] = torch.tensor([]) - for i, k in enumerate(self.prognostic_var_names): - if kt in k: - self.VAR_3D_IDX[kt] = torch.cat( - [self.VAR_3D_IDX[kt], torch.tensor([i])] - ) - self.VAR_3D_IDX[kt] = self.VAR_3D_IDX[kt].to(torch.int32) - - def _populate_dp_3d_idx(self): - for d in self.DEPTH_SET: - self.DP_3D_IDX[d] = torch.tensor([]) - for i, k in enumerate(self.prognostic_var_names): - k_split = k.split("_") - if len(k_split) == 1: - continue - elif d == k_split[-1]: - self.DP_3D_IDX[d] = torch.cat( - [self.DP_3D_IDX[d], torch.tensor([i])] - ) - self.DP_3D_IDX[d] = self.DP_3D_IDX[d].to(torch.int32) - - self.DP_3D_IDX[self.DEPTH_SET[0]] = torch.cat( - [ - self.DP_3D_IDX[self.DEPTH_SET[0]], - torch.tensor([self.VAR_3D_IDX[var_2D] for var_2D in self.VAR_SET_2D]), - ] - ).to(torch.int32) - - def _populate_boundary_idx(self): - """ - Populates the indices of the boundary variables in the input tensor. - - We assume the indices INPT_BOUNDARY_IDX will be used after the boundary - condition is extracted from the input tensor - """ - for i, k in enumerate(self.boundary_var_names): - self.INPT_BOUNDARY_IDX[k] = torch.tensor([i]) - - def to(self, device: torch.device) -> Self: - """Move all index tensors to the given device. - - Call this once after model initialization so that indexing a GPU tensor - with these indices stays on-device and avoids implicit CUDA syncs. - """ - for k in self.VAR_3D_IDX: - self.VAR_3D_IDX[k] = self.VAR_3D_IDX[k].to(device) - for k in self.DP_3D_IDX: - self.DP_3D_IDX[k] = self.DP_3D_IDX[k].to(device) - for k in self.INPT_BOUNDARY_IDX: - self.INPT_BOUNDARY_IDX[k] = self.INPT_BOUNDARY_IDX[k].to(device) - self.dz = self.dz.to(device) - return self diff --git a/src/samudra/datasets.py b/src/samudra/datasets.py index 5955ca41..93201c3b 100644 --- a/src/samudra/datasets.py +++ b/src/samudra/datasets.py @@ -4,31 +4,25 @@ import logging import time -from concurrent.futures import wait from concurrent.futures.thread import ThreadPoolExecutor from typing import ClassVar, final import numpy as np import torch import xarray as xr -from einops import rearrange from jaxtyping import Float from torch.utils.data import Dataset -from xarray_einstats.einops import rearrange as xr_rearrange # noqa: F401 from samudra.constants import ( Boundary, BoundaryVarNames, - Example, - GridMask, - Input, LoaderVersion, Prognostic, - PrognosticMask, PrognosticVarNames, + RolloutStep, ) -from samudra.utils.ctx import GridContext -from samudra.utils.data import DataSource, LoadStats, OceanData, conditional_rearrange +from samudra.utils.ctx import BatchGrid +from samudra.utils.data import BatchPreprocessor, CanonicalSource, LoadStats from samudra.utils.device import using_gpu from samudra.utils.logging import elapsed @@ -52,7 +46,7 @@ class InferenceDataset(Dataset): @elapsed def __init__( self, - src: DataSource, + source: CanonicalSource, prognostic_var_names, boundary_var_names, hist, @@ -66,17 +60,23 @@ def __init__( # using the dataset for inference. self.hist = hist + self.prognostic_var_names = tuple(prognostic_var_names) + self.boundary_var_names = tuple(boundary_var_names) self.num_prognostic_channels = (hist + 1) * len(prognostic_var_names) - data = src.data - self.input_res = src.resolution - self._prognostic_src = src.filter(prognostic_var_names, prefix="prognostic") - self._boundary_src = src.filter(boundary_var_names, prefix="boundary") - self._times = data.time - self.normalize_before_mask = normalize_before_mask - self.masked_fill_value = masked_fill_value - - time_indices = np.arange(data.time.size) + self.input_resolution = source.resolution + self._device = torch.device("cpu") + self._source = source + self._times = source.time + self.preprocessor = BatchPreprocessor( + source, + self.prognostic_var_names, + self.boundary_var_names, + normalize_before_mask=normalize_before_mask, + masked_fill_value=masked_fill_value, + ) + + time_indices = np.arange(source.time.size) indices = xr.DataArray( time_indices, dims=["time"], @@ -96,23 +96,21 @@ def __init__( if long_rollout: logger.info( - f"Long rollout will use input at time {data.time.values[0]} and produce" - f" output at {data.time.values[self.hist + 1]}" + f"Long rollout will use input at time {source.time.values[0]} and produce" + f" output at {source.time.values[self.hist + 1]}" ) - self.wet: PrognosticMask = src.masks.prognostic - self.wet_surface: GridMask = src.masks.boundary - self.wet_label = src.masks.prognostic_with_hist(self.hist) + self.wet_label = source.masks.prognostic_with_hist(self.hist) self.size = len(self.rolling_indices) if using_gpu(): - self.wet = self.wet.pin_memory() - self.wet_surface = self.wet_surface.pin_memory() self.wet_label = self.wet_label.pin_memory() # Inference only currently supports the same output resolution as the input # resolution. - self.ctx = GridContext(self.wet_label, self.input_res, self.input_res) + self.ctx = BatchGrid( + self.wet_label, self.input_resolution, self.input_resolution + ) def __len__(self): return self.size @@ -124,6 +122,7 @@ def to(self, device: torch.device) -> "InferenceDataset": are on the correct device (GPU). """ self.ctx = self.ctx.to(device) + self._device = device self.wet_label = self.wet_label.to(device, non_blocking=True) return self @@ -197,39 +196,11 @@ def _get_x_index(self, idx): return x_index def _get_prognostic(self, x_index): - data_in_src = self._prognostic_src.map_data( - lambda ds: ds.isel(time=x_index).isel(time=slice(None, self.hist + 1)) + return self._read_and_prepare( + x_index.values[:, : self.hist + 1], + channels=self.prognostic_var_names, + prognostic=True, ) - if self.normalize_before_mask: - data_in_ds = data_in_src.normalize() - else: - data_in_ds = data_in_src.data - - if "lev" in data_in_ds.dims: - data_in_np: np.ndarray = ( - conditional_rearrange( - data_in_ds, - "window_dim time (variable lev)=var lat lon", - concat_dim="var", - ) - .rename({"var": "variable"}) - .to_numpy() - ) - else: - data_in_np = ( - data_in_ds.to_array() - .transpose("window_dim", "time", "variable", "lat", "lon") - .to_numpy() - ) - data_in: torch.Tensor = torch.from_numpy(data_in_np).float() - data_in = torch.where(self.wet, data_in, self.masked_fill_value) - if not self.normalize_before_mask: - data_in = self._prognostic_src.normalize_with(data_in, variable_axis=2) - data_in = rearrange( - data_in, - "window_dim time variable lat lon -> window_dim (time variable) lat lon", - ) - return data_in def _get_boundary(self, x_index): """ @@ -238,70 +209,36 @@ def _get_boundary(self, x_index): With hist > 0, the boundary condition considered is always the last step of the input. """ - data_in_boundary_src = self._boundary_src.map_data( - lambda ds: ds.isel(time=x_index).isel(time=slice(None, self.hist + 1)) - ) - if self.normalize_before_mask: - data_in_boundary_ds = data_in_boundary_src.normalize() - else: - data_in_boundary_ds = data_in_boundary_src.data - data_in_boundary_np: np.ndarray = ( - data_in_boundary_ds.to_array() - .transpose("window_dim", "time", "variable", "lat", "lon") - .to_numpy() - ) - data_in_boundary: torch.Tensor = torch.from_numpy(data_in_boundary_np).float() - data_in_boundary = torch.where( - self.wet_surface, data_in_boundary, self.masked_fill_value - ) - if not self.normalize_before_mask: - data_in_boundary = self._boundary_src.normalize_with( - data_in_boundary, variable_axis=2 - ) - data_in_boundary = rearrange( - data_in_boundary, - "window_dim time variable lat lon -> window_dim (time variable) lat lon", + return self._read_and_prepare( + x_index.values[:, : self.hist + 1], + channels=self.boundary_var_names, + prognostic=False, ) - return data_in_boundary def _get_label(self, x_index): - label_src = self._prognostic_src.map_data( - lambda ds: ds.isel(time=x_index).isel(time=slice(self.hist + 1, None)) + return self._read_and_prepare( + x_index.values[:, self.hist + 1 :], + channels=self.prognostic_var_names, + prognostic=True, ) - if self.normalize_before_mask: - label_ds = label_src.normalize() - else: - label_ds = label_src.data - if "lev" in label_ds.dims: - label_np: np.ndarray = ( - conditional_rearrange( - label_ds, - "window_dim time (variable lev)=var lat lon", - concat_dim="var", - ) - .rename({"var": "variable"}) - .to_numpy() - ) - else: - label_np = ( - label_ds.to_array() - .transpose("window_dim", "time", "variable", "lat", "lon") - .to_numpy() - ) - label: torch.Tensor = torch.from_numpy(label_np).float() - label = torch.where(self.wet, label, self.masked_fill_value) - if not self.normalize_before_mask: - label = self._prognostic_src.normalize_with(label, variable_axis=2) - label = rearrange( - label, - "window_dim time variable lat lon -> window_dim (time variable) lat lon", + + def _read_and_prepare( + self, + time_indices: np.ndarray, + *, + channels: tuple[str, ...], + prognostic: bool, + ) -> torch.Tensor: + data = torch.from_numpy(self._source.read(time_indices, channels)) + prepare = ( + self.preprocessor.prepare_prognostic + if prognostic + else self.preprocessor.prepare_boundary ) - return label + return prepare(data, self._device) def get_coords_dict(self): - return { - co: self._prognostic_src.data[co] for co in self._prognostic_src.data.coords - } + return self._source.coordinates() class InferenceDatasets(Dataset): @@ -316,109 +253,76 @@ def __getitem__(self, idx): return (self.datasets[idx], self.lengths[idx]) -class RawTrainData: +class HostBatch: def __init__(self, dataset_id: "TorchTrainDataset.Id"): self.dataset_id: TorchTrainDataset.Id = dataset_id - self.raw_data: list[tuple[torch.Tensor, torch.Tensor, torch.Tensor]] = [] + self.steps: list[RolloutStep] = [] self.load_stats: LoadStats | None = None - def insert( + def append( self, input_: torch.Tensor, boundary: torch.Tensor, label: torch.Tensor, ): """Add a prognostic input, boundary, and prognostic label as the last step.""" - self.raw_data.append((input_, boundary, label)) - - def to(self, device: torch.device): - self.raw_data = [ - ( - input_.to(device, non_blocking=True), - boundary.to(device, non_blocking=True), - label.to(device, non_blocking=True), - ) - for input_, boundary, label in self.raw_data - ] + self.steps.append(RolloutStep(input_, boundary, label)) def pin_memory(self): - self.raw_data = [ - ( + self.steps = [ + RolloutStep( input_.pin_memory(), boundary.pin_memory(), label.pin_memory(), ) - for input_, boundary, label in self.raw_data + for input_, boundary, label in self.steps ] return self -class TrainData: +class ModelBatch: """A single batch of training data. - A single batch contains multiple steps worth of ``Example`` entries, each + A single batch contains multiple steps worth of ``RolloutStep`` entries, each of which is a ``(prognostic_input, boundary_input, label)`` triple. The prognostic and boundary tensors are carried separately because the samudra-multi model encodes them separately (Samudra just concatenates them later). """ - def __init__( - self, num_prognostic_channels: int, num_boundary_channels: int, ctx: GridContext - ): - self.num_prognostic_channels = num_prognostic_channels - self.num_boundary_channels = num_boundary_channels + def __init__(self, ctx: BatchGrid): self.ctx = ctx - self.example_by_step: list[Example] = [] + self.steps: list[RolloutStep] = [] self.load_stats: LoadStats | None = None def append( self, prognostic_input: Prognostic, boundary_input: Boundary, label: Prognostic ) -> None: - """Add another Example as a new step.""" - self.example_by_step.append((prognostic_input, boundary_input, label)) + """Add another RolloutStep as a new step.""" + self.steps.append(RolloutStep(prognostic_input, boundary_input, label)) def get_initial_input(self) -> tuple[Prognostic, Boundary]: return self.get_input(0) def get_input(self, step: int) -> tuple[Prognostic, Boundary]: - prog, boundary, _ = self.example_by_step[step] + prog, boundary, _ = self.steps[step] return prog, boundary def get_label(self, step: int) -> Prognostic: - return self.example_by_step[step][2] + return self.steps[step][2] - def __getitem__(self, step: int) -> Example: + def __getitem__(self, step: int) -> RolloutStep: """Converts index (step) into (prognostic, boundary, label) triple.""" - return self.example_by_step[step] + return self.steps[step] def __len__(self) -> int: - return len(self.example_by_step) + return len(self.steps) def __iter__(self): - return iter(range(len(self))) - - def to(self, device: torch.device) -> None: - for step in self: - prog, boundary, label = self.example_by_step[step] - self.example_by_step[step] = ( - prog.to(device, non_blocking=True), - boundary.to(device, non_blocking=True), - label.to(device, non_blocking=True), - ) - - def pin_memory(self): - for step in self: - prog, boundary, label = self.example_by_step[step] - self.example_by_step[step] = ( - prog.pin_memory(), - boundary.pin_memory(), - label.pin_memory(), - ) - return self + return iter(self.steps) @final -class TorchTrainDataset(Dataset[RawTrainData]): +class TorchTrainDataset(Dataset[HostBatch]): """ This class is used for training and validation. @@ -426,13 +330,13 @@ class TorchTrainDataset(Dataset[RawTrainData]): from InferenceDataset, as it creates rolling indices based on stride. By default, the sliding window / stride is 1. - We make use of TrainData class to store a single sample. + We make use of ModelBatch class to store a single sample. For example, - Hist=0 ; TD: step=0->[0, 1]; step=1->[1, 2]; step=2->[2, 3]; step=3->[3, 4] - Hist=1 ; TD: step=0->[[0, 1], [2, 3]]; step=1->[[2, 3], [4, 5]]; + Hist=0 ; step=0->[0, 1]; step=1->[1, 2]; step=2->[2, 3]; step=3->[3, 4] + Hist=1 ; step=0->[[0, 1], [2, 3]]; step=1->[[2, 3], [4, 5]]; step=2->[[4, 5], [6, 7]]; step=3->[[6, 7], [8, 9]] - Hist=2 ; TD: step=0->[[0, 1, 2], [3, 4, 5]]; + Hist=2 ; step=0->[[0, 1, 2], [3, 4, 5]]; step=1->[[3, 4, 5], [6, 7, 8]]; step=2->[[6, 7, 8], [9, 10, 11]]; step=3->[[9, 10, 11], [12, 13, 14]] @@ -458,7 +362,8 @@ def _get_executor(cls) -> ThreadPoolExecutor: @elapsed def __init__( self, - src: DataSource, + input_source: CanonicalSource, + label_source: CanonicalSource | None, prognostic_var_names: PrognosticVarNames, boundary_var_names: BoundaryVarNames, hist: int, @@ -470,19 +375,20 @@ def __init__( ): super().__init__() self.id = f"{self.__class__.__name__}_{str(id(self))}" + sources = [input_source, label_source] if label_source else [input_source] self.hist: int = hist self.steps: int = steps self.stride: int = stride - self.normalize_before_mask: bool = normalize_before_mask - self.masked_fill_value: float = masked_fill_value self._concurrent_compute = concurrent_compute_ - self.num_prognostic_channels: int = (hist + 1) * len(prognostic_var_names) - self.num_boundary_channels: int = (hist + 1) * len(boundary_var_names) - time_ = src.data.time - self.prognostic_src = src.filter(prognostic_var_names, prefix="prog") - self.boundary_src = src.filter(boundary_var_names, prefix="boundary") + assert np.array_equal(sources[0].time, sources[-1].time), ( + "Input and label sources have different time slices!" + ) + time_ = input_source.time + self.sources = sources + self.prognostic_var_names = tuple(prognostic_var_names) + self.boundary_var_names = tuple(boundary_var_names) # This class will be used only for training and validation total_steps: int = 2 * self.hist + 2 @@ -502,14 +408,21 @@ def __init__( indices_da + stride * window_dim ) - # NB(alxmrs): Keep masks on CPU - will be moved to GPU in to_train_data() - self.wet_prognostic: PrognosticMask = src.masks.prognostic - self.wet_surface: GridMask = src.masks.boundary + self.preprocessors = [ + BatchPreprocessor( + source, + self.prognostic_var_names, + self.boundary_var_names, + normalize_before_mask=normalize_before_mask, + masked_fill_value=masked_fill_value, + ) + for source in sources + ] - self.ctx = GridContext( - label_mask=self.prognostic_src.masks.prognostic_with_hist(self.hist), - input_resolution_cpu=self.prognostic_src.resolution, - output_resolution_cpu=self.prognostic_src.resolution, + self.ctx = BatchGrid( + label_mask=self.sources[-1].masks.prognostic_with_hist(self.hist), + input_resolution_cpu=self.sources[0].resolution, + output_resolution_cpu=self.sources[-1].resolution, ) self.size: int = ( @@ -524,118 +437,66 @@ def __len__(self) -> int: @elapsed(level=logging.DEBUG) def __getitem__(self, idx: int): start_time = time.perf_counter() - TD = RawTrainData(self.id) + host_batch = HostBatch(self.id) for step in range(self.steps): x_index = self._get_x_index(idx, step) current_x_index = x_index.isel(time=slice(0, self.hist + 1)) forecast_x_index = x_index.isel(time=slice(self.hist + 1, None)) - # Only materialize the time ranges we actually use to reduce memory. - input_selected = self.prognostic_src.data.isel(time=current_x_index) - boundary_selected = self.boundary_src.data.isel(time=current_x_index) - label_selected = self.prognostic_src.data.isel( - time=forecast_x_index - ) # forecasted data - prognostic_selected = [input_selected, label_selected] - + reads = ( + ( + self.sources[0], + current_x_index.values, + self.prognostic_var_names, + ), + ( + self.sources[0], + current_x_index.values, + self.boundary_var_names, + ), + ( + self.sources[-1], + forecast_x_index.values, + self.prognostic_var_names, + ), + ) if self._concurrent_compute: - datasets = prognostic_selected + [boundary_selected] - concurrent_compute( - *datasets, - executor=self._get_executor(), - ) - - if "lev" in prognostic_selected[0].dims: - prognostics = [ - torch.from_numpy( - conditional_rearrange( - selected, - "time (variable lev)=var lat lon", - concat_dim="var", - ) - .rename({"var": "variable"}) - .to_numpy() - .astype(np.float32, copy=False) - ) - for selected in prognostic_selected + executor = self._get_executor() + futures = [ + executor.submit(source.read, indices, channels) + for source, indices, channels in reads ] + loaded = [future.result() for future in futures] else: - prognostics = [ - torch.from_numpy( - selected.to_array() - .transpose("time", "variable", "lat", "lon") - .to_numpy() - .astype(np.float32, copy=False) - ) - for selected in prognostic_selected + loaded = [ + source.read(indices, channels) + for source, indices, channels in reads ] - boundary = torch.from_numpy( - boundary_selected.to_array() - .transpose("time", "variable", "lat", "lon") - .to_numpy() - .astype(np.float32, copy=False) - ) - input_, label = prognostics[0], prognostics[-1] - TD.insert(input_, boundary, label) - TD.load_stats = LoadStats(time.perf_counter() - start_time) + input_, boundary, label = map(torch.from_numpy, loaded) + host_batch.append(input_, boundary, label) + host_batch.load_stats = LoadStats(time.perf_counter() - start_time) - return TD + return host_batch - def to_train_data( - self, raw_train_data: RawTrainData, device: torch.device - ) -> TrainData: - """Convert RawTrainData to TrainData, moving tensors to the specified device. + def to_model_batch(self, host_batch: HostBatch, device: torch.device) -> ModelBatch: + """Convert HostBatch to ModelBatch, moving tensors to the specified device. Args: - raw_train_data: CPU data from worker process + host_batch: CPU data from worker process device: Target device (typically GPU) to move tensors to Returns: - TrainData with tensors on the target device + ModelBatch with tensors on the target device """ - train_data = TrainData( - self.num_prognostic_channels, - self.num_boundary_channels, - self.ctx.to(device), - ) - for input_, boundary, label in raw_train_data.raw_data: - prog_input, boundary_input, label_tensor = self._to_example( - OceanData.from_data_source( - input_, - self.wet_prognostic, - self.prognostic_src, - ).to(device=device, non_blocking=True), - OceanData.from_data_source( - boundary, - self.wet_surface, - self.boundary_src, - ).to(device=device, non_blocking=True), - OceanData.from_data_source( - label, self.wet_prognostic, self.prognostic_src - ).to(device=device, non_blocking=True), - ) - train_data.append(prog_input, boundary_input, label_tensor) - train_data.load_stats = raw_train_data.load_stats - return train_data - - def _to_example( - self, input_: OceanData, boundary: OceanData, label: OceanData - ) -> Example: - # Input/boundary only include current steps; label only includes forecasted steps. - prog_input = self._prep_tensor_steps(input_) - boundary_input = self._prep_tensor_steps(boundary) - label_tensor = self._prep_tensor_steps(label) - return prog_input, boundary_input, label_tensor - - def _prep_tensor_steps(self, ocean_data: OceanData) -> Input: - """Normalize, mask, and flatten (time, variable) dims into a channel dim.""" - steps = ocean_data.normalize_and_mask( - self.normalize_before_mask, self.masked_fill_value - ) - return rearrange( - steps, "batch time variable lat lon -> batch (time variable) lat lon" - ) + model_batch = ModelBatch(self.ctx.to(device)) + for input_, boundary, label in host_batch.steps: + prog_input = self.preprocessors[0].prepare_prognostic(input_, device) + boundary_input = self.preprocessors[0].prepare_boundary(boundary, device) + label_tensor = self.preprocessors[-1].prepare_prognostic(label, device) + model_batch.append(prog_input, boundary_input, label_tensor) + model_batch.load_stats = host_batch.load_stats + return model_batch def _get_x_index(self, idx: int, step: int) -> xr.DataArray: assert isinstance(idx, int) @@ -648,27 +509,12 @@ def _get_x_index(self, idx: int, step: int) -> xr.DataArray: return self.rolling_indices.isel(window=window_index, drop=True) -def concurrent_compute( - *datasets: xr.Dataset, - executor: ThreadPoolExecutor, -) -> None: - def load_variable_data(var: xr.Variable) -> None: - var.load() - - futures = [] - for ds in datasets: - for var in ds.data_vars.variables.values(): - futures.append(executor.submit(load_variable_data, var)) - - wait(futures) - - @final -class TrainDataLoader: +class BatchLoader: """Wrapper around a torch DataLoader that handles GPU post-processing. - This class wraps a DataLoader[RawTrainData] and converts the raw data - to TrainData by applying GPU-based normalization and masking. This allows + This class wraps a DataLoader[HostBatch] and converts the raw data + to ModelBatch by applying GPU-based normalization and masking. This allows the data loading process to handle I/O while the main process handles GPU operations. @@ -681,46 +527,46 @@ class TrainDataLoader: def __init__( self, - dataloader: torch.utils.data.DataLoader[RawTrainData], + host_loader: torch.utils.data.DataLoader[HostBatch], datasets: list[TorchTrainDataset], device: torch.device, ): - self._dataloader = dataloader + self._host_loader = host_loader self._datasets = {dataset.id: dataset for dataset in datasets} self._device = device def __iter__(self): - """Iterate over the dataloader, converting RawTrainData to TrainData.""" - for raw_train_data in self._dataloader: - dataset = self._datasets[raw_train_data.dataset_id] - train_data = dataset.to_train_data(raw_train_data, self._device) - yield train_data + """Iterate over the dataloader, converting HostBatch to ModelBatch.""" + for host_batch in self._host_loader: + dataset = self._datasets[host_batch.dataset_id] + model_batch = dataset.to_model_batch(host_batch, self._device) + yield model_batch def __len__(self) -> int: - return len(self._dataloader) + return len(self._host_loader) - def __getitem__(self, index: int) -> TrainData: - """Access a single item by index, converting RawTrainData to TrainData. + def __getitem__(self, index: int) -> ModelBatch: + """Access a single item by index, converting HostBatch to ModelBatch. Note: This bypasses the DataLoader's sampling/batching and directly accesses the underlying dataset for test purposes. """ # Access the underlying dataset directly - raw_train_data = self._dataloader.dataset[index] + host_batch = self._host_loader.dataset[index] # Apply the collate function to add batch dimension (expects a list) - collate_fn = self._dataloader.collate_fn + collate_fn = self._host_loader.collate_fn if collate_fn is not None: - raw_train_data = collate_fn([raw_train_data]) + host_batch = collate_fn([host_batch]) # Get the dataset that created this raw data - dataset = self._datasets[raw_train_data.dataset_id] - # Convert to TrainData - train_data = dataset.to_train_data(raw_train_data, self._device) - return train_data + dataset = self._datasets[host_batch.dataset_id] + # Convert to ModelBatch + model_batch = dataset.to_model_batch(host_batch, self._device) + return model_batch @property def dataset(self): - return self._dataloader.dataset + return self._host_loader.dataset @property def sampler(self): - return self._dataloader.sampler + return self._host_loader.sampler diff --git a/src/samudra/derived_variables.py b/src/samudra/derived_variables.py index ae5d6f8f..da7697e8 100644 --- a/src/samudra/derived_variables.py +++ b/src/samudra/derived_variables.py @@ -7,7 +7,7 @@ import torch from jaxtyping import Float -from samudra.constants import CP_SW, RHO_0, Grid, TensorMap +from samudra.constants import CP_SW, RHO_0, DataLayout, Grid def compute_ocean_heat_content( @@ -74,16 +74,16 @@ def compute_global_ocean_heat_content( def add_derived_variables( tensor_out: torch.Tensor, *, - tensor_map: TensorMap, + data_layout: DataLayout, ) -> dict[str, torch.Tensor]: """ Add derived variables to the output. """ # Ocean heat content derived_vars = {} - dz = tensor_map.dz.to(tensor_out.device) + dz = data_layout.dz.to(tensor_out.device) thetao = tensor_out[ - :, :, tensor_map.VAR_3D_IDX[tensor_map.dataset_spec.ocean_heat_temperature_var] + :, :, data_layout.variable_indices[data_layout.ocean_heat_temperature_var] ] ohct = compute_ocean_heat_content(thetao, dz) derived_vars["ocean_heat_content"] = ohct diff --git a/src/samudra/eval.py b/src/samudra/eval.py index 9d473c35..fc573845 100644 --- a/src/samudra/eval.py +++ b/src/samudra/eval.py @@ -12,10 +12,10 @@ from samudra.aggregator import Aggregator from samudra.backend import init_eval_backend from samudra.config import EvalConfig -from samudra.constants import BoundaryVarNames, Grid, PrognosticVarNames, TensorMap +from samudra.constants import BoundaryVarNames, Grid, PrognosticVarNames from samudra.datasets import InferenceDataset from samudra.stepper import run_rollout -from samudra.utils.data import Normalize, get_inference_steps, spherical_area_weights +from samudra.utils.data import BatchPreprocessor, get_inference_steps from samudra.utils.device import using_gpu from samudra.utils.distributed import is_main_process, set_seed from samudra.utils.logging import get_model_summary, handle_logging, handle_warnings @@ -47,17 +47,17 @@ def __init__(self, cfg: EvalConfig) -> None: set_seed(cfg.experiment.rand_seed) logger.info("Loading data") - self.data_container = cfg.data.build( + self.data_bundle = cfg.data.build( cfg.experiment.resolved_data_root, ) # Getting prognostic and boundary variables - self.dataset_spec = self.data_container.dataset_spec + self.data_layout = self.data_bundle.data_layout self.prognostic_var_names: PrognosticVarNames = ( - self.dataset_spec.prognostic_var_names + self.data_layout.prognostic_var_names ) - self.boundary_var_names: BoundaryVarNames = self.dataset_spec.boundary_var_names - self.levels = self.dataset_spec.num_prognostic_depth_levels + self.boundary_var_names: BoundaryVarNames = self.data_layout.boundary_var_names + self.levels = self.data_layout.num_prognostic_depth_levels str_prognostics = ", ".join([i for i in self.prognostic_var_names]) str_boundaries = ", ".join([i for i in self.boundary_var_names]) @@ -74,25 +74,24 @@ def __init__(self, cfg: EvalConfig) -> None: self.num_in = self.num_prog_in + self.num_boundary_in self.num_out = self.num_prog_in - self.tensor_map = TensorMap(dataset_spec=self.dataset_spec).to(self.device) + self.data_layout = self.data_layout.to(self.device) logger.info(f"Number of inputs (prognostic + boundary): {self.num_in}") logger.info(f"Number of outputs (prognostic): {self.num_out}") # Dataloaders - if self.data_container.inference_source is None: + if self.data_bundle.inference_source is None: raise ValueError( "Inference time is not configured for the first data source" ) - self.src = self.data_container.inference_source - self.data = self.src.data - self.metadata = self.src.metadata - self.wet = self.src.masks.prognostic_with_hist(cfg.data.hist) - self.area_weights: Grid = spherical_area_weights(self.data) + self.source = self.data_bundle.inference_source + self.metadata = self.source.metadata + self.wet = self.source.masks.prognostic_with_hist(cfg.data.hist) + self.area_weights: Grid = self.source.spherical_area_weights self.area_weights = self.area_weights.to(self.device) - self.normalize = Normalize( - self.src, + self.preprocessor = BatchPreprocessor( + self.source, prognostic_var_names=self.prognostic_var_names, boundary_var_names=self.boundary_var_names, ) @@ -103,7 +102,7 @@ def __init__(self, cfg: EvalConfig) -> None: boundary_channels=self.num_boundary_in, out_channels=self.num_out, hist=cfg.data.hist, - srcs=self.data_container.train_sources, + grid_sizes=[source.grid_size for source in self.data_bundle.train_sources], ).to(self.device) get_model_summary(self.model, None, cfg.debug) @@ -124,7 +123,7 @@ def __init__(self, cfg: EvalConfig) -> None: # Set up wandb run self.wandb_id, self.wandb_name = self.wandb_logger.setup_run( - None, cfg, data_container=self.data_container, finetune=False + None, cfg, data_bundle=self.data_bundle, finetune=False ) # Eval @@ -150,11 +149,11 @@ def load_checkpoint(self, ckpt_path: str): def init_inference_store(self): self.num_time_steps = get_inference_steps( - self.src, + self.source, hist=self.hist, ) self.inference_dataset = InferenceDataset( - src=self.src, + source=self.source, prognostic_var_names=self.prognostic_var_names, boundary_var_names=self.boundary_var_names, hist=self.hist, @@ -189,10 +188,10 @@ def standalone_inference(self): self.metadata, self.hist, self.area_weights, - self.src.masks.prognostic.to(self.device), + self.source.masks.prognostic.to(self.device), self.num_out, - self.tensor_map, - self.normalize, + self.data_layout, + self.preprocessor, self.prognostic_var_names, ) @@ -205,8 +204,8 @@ def standalone_inference(self): model_path=self.model_path, num_model_steps_forward=self.num_model_steps_forward, save_zarr=self.save_zarr, - tensor_map=self.tensor_map, - normalize=self.normalize, + data_layout=self.data_layout, + preprocessor=self.preprocessor, ) logs = inf_aggregator.get_summary_logs() return {f"inference/{k}": v for k, v in logs.items()} diff --git a/src/samudra/models/base.py b/src/samudra/models/base.py index e0ee5329..f8332f21 100644 --- a/src/samudra/models/base.py +++ b/src/samudra/models/base.py @@ -9,11 +9,11 @@ import torch from samudra.constants import Boundary, Prognostic -from samudra.utils.ctx import GridContext +from samudra.utils.ctx import BatchGrid logger = logging.getLogger(__name__) -from samudra.datasets import InferenceDataset, TrainData +from samudra.datasets import InferenceDataset, ModelBatch from samudra.utils.device import get_device from samudra.utils.output import ModelInferenceOutput @@ -47,20 +47,20 @@ def __init__( self.gradient_detach_interval = gradient_detach_interval def forward_once( - self, prognostic: Prognostic, boundary: Boundary, ctx: GridContext + self, prognostic: Prognostic, boundary: Boundary, ctx: BatchGrid ) -> Prognostic: raise NotImplementedError() def forward( self, - train_data: TrainData, + batch: ModelBatch, loss_fn=None, ) -> torch.Tensor | list[torch.Tensor]: outputs: list[torch.Tensor] = [] loss = torch.tensor(torch.nan) - for step in range(len(train_data)): + for step in range(len(batch)): if step == 0: - prog_tensor, boundary_tensor = train_data.get_initial_input() + prog_tensor, boundary_tensor = batch.get_initial_input() else: prev_output = outputs[-1] if ( @@ -68,10 +68,10 @@ def forward( and step % self.gradient_detach_interval == 0 ): prev_output = prev_output.detach() - _, boundary_tensor = train_data.get_input(step) + _, boundary_tensor = batch.get_input(step) prog_tensor = prev_output - decodings = self.forward_once(prog_tensor, boundary_tensor, train_data.ctx) + decodings = self.forward_once(prog_tensor, boundary_tensor, batch.ctx) if self.pred_residuals: pred = prog_tensor + decodings # Residual prediction else: @@ -81,12 +81,12 @@ def forward( if step == 0: loss = loss_fn( pred, - train_data.get_label(step), + batch.get_label(step), ) else: loss += loss_fn( pred, - train_data.get_label(step), + batch.get_label(step), ) outputs.append(pred) @@ -139,5 +139,5 @@ def inference( slice(steps_completed, steps_completed + num_steps) ).to(device=get_device()) - IO = ModelInferenceOutput(pred_tensor, target_tensor, target_time) - return IO + inference_output = ModelInferenceOutput(pred_tensor, target_tensor, target_time) + return inference_output diff --git a/src/samudra/models/samudra.py b/src/samudra/models/samudra.py index 75763cd5..c7341af3 100644 --- a/src/samudra/models/samudra.py +++ b/src/samudra/models/samudra.py @@ -9,7 +9,7 @@ from samudra.constants import Boundary, GridSize, Prognostic from samudra.models.base import BaseModel from samudra.models.modules.unet_backbone import UNetBackbone -from samudra.utils.ctx import GridContext +from samudra.utils.ctx import BatchGrid from samudra.utils.device import autocast @@ -58,7 +58,7 @@ def __init__( self.use_bfloat16 = use_bfloat16 def forward_once( - self, prognostic: Prognostic, boundary: Boundary, ctx: GridContext + self, prognostic: Prognostic, boundary: Boundary, ctx: BatchGrid ) -> Prognostic: # Samudra is a single-scale model; fuse prognostic + boundary into # the single channel-stacked input its backbone expects. diff --git a/src/samudra/models/samudra_mini.py b/src/samudra/models/samudra_mini.py index d0e625a9..8cd128e1 100644 --- a/src/samudra/models/samudra_mini.py +++ b/src/samudra/models/samudra_mini.py @@ -16,7 +16,7 @@ from samudra.constants import Boundary, Prognostic from samudra.models.base import BaseModel from samudra.models.modules.augment_input import make_3d_coordinate_grid -from samudra.utils.ctx import GridContext +from samudra.utils.ctx import BatchGrid from samudra.utils.device import autocast if TYPE_CHECKING: @@ -131,7 +131,7 @@ def _decode( return torch.cat(out_chunks, dim=1) def forward_once( - self, prognostic: Prognostic, boundary: Boundary, ctx: GridContext + self, prognostic: Prognostic, boundary: Boundary, ctx: BatchGrid ) -> Prognostic: # SamudraMini is a single-scale pixel-token model; fuse prognostic + # boundary into the single channel-stacked input it expects. diff --git a/src/samudra/models/samudra_multi.py b/src/samudra/models/samudra_multi.py index dbcca07e..07b12e7c 100644 --- a/src/samudra/models/samudra_multi.py +++ b/src/samudra/models/samudra_multi.py @@ -16,7 +16,7 @@ from samudra.models.base import BaseModel from samudra.models.modules import PerceiverDecoder, PerceiverEncoder from samudra.models.modules.unet_backbone import UNetBackbone -from samudra.utils.ctx import GridContext +from samudra.utils.ctx import BatchGrid from samudra.utils.device import autocast if TYPE_CHECKING: @@ -91,7 +91,7 @@ def __init__( ) def forward_once( - self, prognostic: Prognostic, boundary: Boundary, ctx: GridContext + self, prognostic: Prognostic, boundary: Boundary, ctx: BatchGrid ) -> Prognostic: # Prognostic and boundary are carried as separate tensors through the # data pipeline, but this encoder still expects a single concatenated diff --git a/src/samudra/stepper.py b/src/samudra/stepper.py index d1eaa3d0..296b716a 100644 --- a/src/samudra/stepper.py +++ b/src/samudra/stepper.py @@ -17,10 +17,10 @@ import torch from samudra.aggregator import InferenceEvaluatorAggregator -from samudra.constants import TensorMap -from samudra.datasets import InferenceDataset, TrainData +from samudra.constants import DataLayout +from samudra.datasets import InferenceDataset, ModelBatch from samudra.models.base import BaseModel -from samudra.utils.data import Normalize +from samudra.utils.data import BatchPreprocessor from samudra.utils.device import get_device from samudra.utils.output import ModelInferenceOutput, TrainBatchOutput, ValBatchOutput from samudra.utils.wandb import get_record_to_wandb @@ -30,7 +30,7 @@ def train_batch( - model: torch.nn.Module, batch: TrainData, loss_fn: Callable + model: torch.nn.Module, batch: ModelBatch, loss_fn: Callable ) -> TrainBatchOutput: loss_per_channel = model(batch, loss_fn=partial(loss_fn, ctx=batch.ctx)) loss = torch.mean(loss_per_channel) @@ -40,7 +40,7 @@ def train_batch( @torch.no_grad() def validate_batch( model: BaseModel | torch.nn.parallel.DistributedDataParallel, - batch: TrainData, + batch: ModelBatch, loss_fn: Callable, ) -> ValBatchOutput: assert len(batch) == 1 # Assert we are using one step of input and output @@ -67,8 +67,8 @@ def run_rollout( model_path: str | PathLike | None = None, num_model_steps_forward: int = 200, save_zarr: bool = False, - tensor_map: TensorMap | None = None, - normalize: Normalize | None = None, + data_layout: DataLayout | None = None, + preprocessor: BatchPreprocessor | None = None, ) -> None: """Performs inference, which is an auto-regressive rollout.""" if save_zarr: @@ -76,9 +76,9 @@ def run_rollout( raise ValueError( "output_dir and model_path must be provided if save_zarr is True" ) - if tensor_map is None or normalize is None: + if data_layout is None or preprocessor is None: raise ValueError( - "tensor_map and normalize must be provided if save_zarr is True" + "data_layout and preprocessor must be provided if save_zarr is True" ) coords = dataset.get_coords_dict() if num_model_steps_forward > 0: @@ -91,8 +91,8 @@ def run_rollout( hist=inf_aggregator.hist, model_path=model_path, time_chunk_size=chunk_size, - normalize=normalize, - tensor_map=tensor_map, + preprocessor=preprocessor, + data_layout=data_layout, ) else: writer = None @@ -128,7 +128,7 @@ def run_rollout( f"Stepping {num_steps} steps forward." ) dataset.to(get_device()) - IO: ModelInferenceOutput = model.inference( + inference_output: ModelInferenceOutput = model.inference( dataset, initial_prognostic=initial_prognostic, steps_completed=step, @@ -136,14 +136,14 @@ def run_rollout( epoch=epoch, ) # Setting initial prognostic for next loop - initial_prognostic = IO.prediction[-1].unsqueeze(0).clone() + initial_prognostic = inference_output.prediction[-1].unsqueeze(0).clone() if writer: logger.info("Writing to zarr...") - writer.record_batch(IO) + writer.record_batch(inference_output) writer.write() logger.info("Recording logs...") - logs = inf_aggregator.record_batch(IO) + logs = inf_aggregator.record_batch(inference_output) logger.info("Logging to wandb...") record_logs(logs) step += num_steps diff --git a/src/samudra/train.py b/src/samudra/train.py index 36c0b364..ad4a8b2a 100644 --- a/src/samudra/train.py +++ b/src/samudra/train.py @@ -38,15 +38,14 @@ MAX_TRAIN_MODEL_STEPS_FORWARD, BoundaryVarNames, PrognosticVarNames, - TensorMap, ) from samudra.datasets import ( + BatchLoader, + HostBatch, InferenceDataset, InferenceDatasets, - RawTrainData, + ModelBatch, TorchTrainDataset, - TrainData, - TrainDataLoader, ) from samudra.models.base import BaseModel from samudra.stepper import ( @@ -56,7 +55,7 @@ train_batch, validate_batch, ) -from samudra.utils.data import Normalize, get_inference_steps +from samudra.utils.data import BatchPreprocessor, get_inference_steps from samudra.utils.device import using_gpu from samudra.utils.distributed import ( all_reduce_mean, @@ -79,8 +78,8 @@ ) from samudra.utils.train import ( CheckpointPaths, + collate_host_batches, collate_inference_data, - collate_raw_train_data, ) from samudra.utils.train_progress import TrainProgress from samudra.utils.wandb import WandBLogger @@ -128,17 +127,17 @@ def __init__(self, cfg: TrainConfig) -> None: self.rand_seed = cfg.experiment.rand_seed set_seed(self.rand_seed) - self.data_container = cfg.data.build( + self.data_bundle = cfg.data.build( data_root=cfg.experiment.resolved_data_root, ) # Getting prognostic and boundary variables - self.dataset_spec = self.data_container.dataset_spec + self.data_layout = self.data_bundle.data_layout self.prognostic_var_names: PrognosticVarNames = ( - self.dataset_spec.prognostic_var_names + self.data_layout.prognostic_var_names ) - self.boundary_var_names: BoundaryVarNames = self.dataset_spec.boundary_var_names - self.levels = self.dataset_spec.num_prognostic_depth_levels + self.boundary_var_names: BoundaryVarNames = self.data_layout.boundary_var_names + self.levels = self.data_layout.num_prognostic_depth_levels str_prognostics = ", ".join([i for i in self.prognostic_var_names]) str_boundaries = ", ".join([i for i in self.boundary_var_names]) @@ -162,7 +161,7 @@ def __init__(self, cfg: TrainConfig) -> None: self.num_in = self.num_prog_in + self.num_boundary_in self.num_out = self.num_prog_in - self.tensor_map = TensorMap(dataset_spec=self.dataset_spec).to(self.device) + self.data_layout = self.data_layout.to(self.device) logger.info(f"Number of inputs (prognostic + boundary): {self.num_in}") logger.info(f"Number of outputs (prognostic): {self.num_out}") @@ -178,18 +177,18 @@ def __init__(self, cfg: TrainConfig) -> None: logger.info(f"Loading data") self.concurrent_compute = cfg.data.concurrent_compute - self.primary_src = self.data_container.primary_source + self.primary_source = self.data_bundle.train_sources[0] # We use dask for inference since it has memory issues otherwise. # TODO(jder): Could rewrite inference dataset like we did for TorchTrainDataset # see https://github.com/m2lines/Samudra/issues/208 - self.inference_src = self.data_container.inference_source + self.inference_source = self.data_bundle.inference_source - self.loader_version = self.data_container.loader_version + self.loader_version = self.data_bundle.loader_version # Aggregation still works on the primary source only. - self.normalize = Normalize( - self.primary_src, + self.preprocessor = BatchPreprocessor( + self.primary_source, prognostic_var_names=self.prognostic_var_names, boundary_var_names=self.boundary_var_names, ) @@ -199,7 +198,7 @@ def __init__(self, cfg: TrainConfig) -> None: boundary_channels=self.num_boundary_in, out_channels=self.num_out, hist=cfg.data.hist, - srcs=self.data_container.train_sources, + grid_sizes=[source.grid_size for source in self.data_bundle.train_sources], ).to(self.device) self.nets_dir = cfg.experiment.nets_dir @@ -240,7 +239,7 @@ def __init__(self, cfg: TrainConfig) -> None: self.wandb_id, self.wandb_name = self.wandb_logger.setup_run( cfg.resume_ckpt_path, cfg, - data_container=self.data_container, + data_bundle=self.data_bundle, finetune=cfg.finetune, ) @@ -324,7 +323,7 @@ def __init__(self, cfg: TrainConfig) -> None: self.hist + 1 ) self.normalize_before_mask: bool = cfg.data.normalize_before_mask - self.normalize_fill_value: float = cfg.data.masked_fill_value + self.masked_fill_value: float = cfg.data.masked_fill_value self.delayed_loss_estimate: bool = cfg.delayed_loss_estimate self.profiler = cfg.profiler.build(self.output_dir, self.device) @@ -332,10 +331,10 @@ def __init__(self, cfg: TrainConfig) -> None: self.wandb_logger.enabled ) - assert self.tensor_map is not None + assert self.data_layout is not None if self.inference_epochs: - if self.inference_src is None: + if self.inference_source is None: raise ValueError( "Inference time is not configured for the first data source" ) @@ -351,23 +350,23 @@ def __init__(self, cfg: TrainConfig) -> None: self.inference_sampler: DistributedSampler | RandomSampler # Add type annotations for loaders - self.train_loader: TrainDataLoader - self.val_loader: TrainDataLoader - self.inference_loader: DataLoader[TrainData] + self.train_loader: BatchLoader + self.val_loader: BatchLoader + self.inference_loader: DataLoader[ModelBatch] def init_inference_stores(self): - assert self.inference_src is not None + assert self.inference_source is not None num_time_steps = get_inference_steps( - self.inference_src, + self.inference_source, hist=self.hist, ) inference_dataset = InferenceDataset( - src=self.inference_src, + source=self.inference_source, prognostic_var_names=self.prognostic_var_names, boundary_var_names=self.boundary_var_names, hist=self.hist, normalize_before_mask=self.normalize_before_mask, - masked_fill_value=self.normalize_fill_value, + masked_fill_value=self.masked_fill_value, long_rollout=True, ) @@ -466,7 +465,7 @@ def run(self) -> None: def train_one_epoch(self, epoch): self.model.train(True) - train_aggregator = Aggregator.get_train_aggregator(self.tensor_map) + train_aggregator = Aggregator.get_train_aggregator(self.data_layout) metric_logger = MetricLogger(delimiter=" ") metric_logger.add_meter("lr", SmoothedValue(window_size=1, fmt="{value:.6f}")) header = f"Training Epoch: [{epoch}]" @@ -485,14 +484,14 @@ def train_one_epoch(self, epoch): else total_batches ) - for data_iter_step, data in enumerate( + for batch_index, batch in enumerate( metric_logger.log_every(self.train_loader, 1, header) ): - if self.debug and (data_iter_step + 1) % 5 == 0: + if self.debug and (batch_index + 1) % 5 == 0: break in_final_cycle = ( - data_iter_step + 1 > final_cycle_start + batch_index + 1 > final_cycle_start ) and remaining_batches > 0 # Determine the actual number of microbatches in this accumulation cycle @@ -502,25 +501,25 @@ def train_one_epoch(self, epoch): r = self.gradient_accumulation_steps if self.num_batches_seen == 0: - get_model_summary(self.model, data, self.debug) + get_model_summary(self.model, batch, self.debug) with self.train_progress.batch( - data, world_size=self.world_size, device=self.device + batch, world_size=self.world_size, device=self.device ) as batch_progress: - TO: TrainBatchOutput = train_batch(self.model, data, self.loss_fn) + batch_output: TrainBatchOutput = train_batch( + self.model, batch, self.loss_fn + ) # Scale loss by this accumulation cycle's actual microbatch count. - scaled_loss = TO.loss / r + scaled_loss = batch_output.loss / r scaled_loss.backward() - train_aggregator.record_batch(TO) + train_aggregator.record_batch(batch_output) self.num_batches_seen += 1 - is_last = data_iter_step + 1 == total_batches - should_step = ( - data_iter_step + 1 - ) % self.gradient_accumulation_steps == 0 + is_last = batch_index + 1 == total_batches + should_step = (batch_index + 1) % self.gradient_accumulation_steps == 0 optimizer_stepped = should_step or is_last batch_progress.optimizer_stepped = optimizer_stepped # Step optimizer after accumulating enough batches or at the end @@ -542,8 +541,10 @@ def train_one_epoch(self, epoch): with torch.no_grad(): # Reduce losses - loss_value_reduce = all_reduce_mean(TO.loss.detach()) - loss_per_channel_reduce = all_reduce_mean(TO.loss_per_channel.detach()) + loss_value_reduce = all_reduce_mean(batch_output.loss.detach()) + loss_per_channel_reduce = all_reduce_mean( + batch_output.loss_per_channel.detach() + ) metrics = { "train/batch/loss": loss_value_reduce, "train/batch/lr": lr, @@ -554,17 +555,17 @@ def train_one_epoch(self, epoch): **get_channel_loss_dict( label="train", loss_per_channel=loss_per_channel_reduce, - tensor_map=self.tensor_map, + data_layout=self.data_layout, ), **get_depth_loss_dict( label="train", loss_per_channel=loss_per_channel_reduce, - tensor_map=self.tensor_map, + data_layout=self.data_layout, ), **get_variable_loss_dict( label="train", loss_per_channel=loss_per_channel_reduce, - tensor_map=self.tensor_map, + data_layout=self.data_layout, ), "train/batch/data_load_time": metric_logger.meters[ "data_load_time" @@ -580,7 +581,7 @@ def train_one_epoch(self, epoch): loss_scale_per_channel = loss_scale_per_channel_fn() # Reshape from time-major channels to [hist, var] and # average along the history dimension. - loss_per_channel = TO.loss_per_channel.reshape( + loss_per_channel = batch_output.loss_per_channel.reshape( -1, loss_scale_per_channel.shape[0] ).mean(dim=0) @@ -594,12 +595,12 @@ def train_one_epoch(self, epoch): **get_channel_loss_scale_dict( label="train", loss_scale_per_channel=loss_scale_per_channel, - tensor_map=self.tensor_map, + data_layout=self.data_layout, ), **get_channel_loss_dict( label="train", loss_per_channel=unscaled_loss_per_channel, - tensor_map=self.tensor_map, + data_layout=self.data_layout, loss_name="loss_unscaled", ), "train/batch/loss_unscaled": unscaled_loss, @@ -614,7 +615,7 @@ def train_one_epoch(self, epoch): metric_logger.update(loss=loss_value_reduce.item()) metric_logger.update(lr=lr) - self._maybe_update_loss(TO, data) + self._maybe_update_loss(batch_output, batch) self.profiler.after_batch(self.num_batches_seen) @@ -624,7 +625,7 @@ def train_one_epoch(self, epoch): logger.info(f"Aggregating train logs") return train_aggregator.get_logs() - def _maybe_update_loss(self, output: TrainBatchOutput, data: TrainData): + def _maybe_update_loss(self, output: TrainBatchOutput, batch: ModelBatch): if (update := getattr(self.loss_fn, "update", None)) is None: return @@ -649,17 +650,15 @@ def _maybe_update_loss(self, output: TrainBatchOutput, data: TrainData): # Run a fresh single-step forward pass so DynamicLoss sees an # up-to-date, unscaled loss signal with torch.no_grad(): - single_step_data = TrainData( - data.num_prognostic_channels, data.num_boundary_channels, data.ctx - ) - prog_input, boundary_input, label = data[0] - single_step_data.append(prog_input, boundary_input, label) - pred = self.model(single_step_data) + single_step_batch = ModelBatch(batch.ctx) + prog_input, boundary_input, label = batch[0] + single_step_batch.append(prog_input, boundary_input, label) + pred = self.model(single_step_batch) # Compute the raw (unscaled) per-channel loss via the inner # loss function, bypassing DynamicLoss scaling. if not isinstance(self.loss_fn, DynamicLoss): raise TypeError(f"Expected loss_fn to be DynamicLoss") - raw_loss = self.loss_fn.loss_fn(pred[0], label, ctx=data.ctx) + raw_loss = self.loss_fn.loss_fn(pred[0], label, ctx=batch.ctx) update(raw_loss) def validate_one_epoch(self, epoch): @@ -670,27 +669,29 @@ def validate_one_epoch(self, epoch): ) val_aggregator = Aggregator.get_validation_aggregator( - self.primary_src.metadata, + self.primary_source.metadata, self.hist, - self.primary_src.spherical_area_weights.to(self.device), + self.primary_source.spherical_area_weights.to(self.device), self.num_out, - self.tensor_map, - self.normalize, + self.data_layout, + self.preprocessor, include_image_aggregators=log_validation_images, ) metric_logger = MetricLogger(delimiter=" ") header = f"One-Step Validation Epoch: [{epoch}]" with torch.no_grad(), self._test_context(): - for data_iter_step, data in enumerate( + for batch_index, batch in enumerate( metric_logger.log_every(self.val_loader, 1, header) ): - if self.debug and (data_iter_step + 1) % 5 == 0: + if self.debug and (batch_index + 1) % 5 == 0: break - VO: ValBatchOutput = validate_batch(self.model, data, self.loss_fn) - val_aggregator.record_validation_batch(VO) - metric_logger.update(loss=VO.loss) + validation_output: ValBatchOutput = validate_batch( + self.model, batch, self.loss_fn + ) + val_aggregator.record_validation_batch(validation_output) + metric_logger.update(loss=validation_output.loss) logger.info(f"Aggregating validation logs") return val_aggregator.get_logs(label="val") @@ -699,19 +700,19 @@ def inference_one_epoch(self, epoch): self.model.eval() with torch.no_grad(), self._test_context(): - for data_iter_step, (inference_dataset, num_steps) in enumerate( + for batch_index, (inference_dataset, num_steps) in enumerate( self.inference_loader ): # TODO(alxmrs): Aggregator only supports a single scale. inf_aggregator = Aggregator.get_inline_inference_aggregator( num_steps, - self.primary_src.metadata, + self.primary_source.metadata, self.hist, - self.primary_src.spherical_area_weights.to(self.device), - self.primary_src.masks.prognostic.to(self.device), + self.primary_source.spherical_area_weights.to(self.device), + self.primary_source.masks.prognostic.to(self.device), self.num_out, - self.tensor_map, - self.normalize, + self.data_layout, + self.preprocessor, self.prognostic_var_names, ) @@ -727,8 +728,8 @@ def inference_one_epoch(self, epoch): num_model_steps_forward=min( num_steps // 2, self.max_train_model_steps_forward ), - tensor_map=self.tensor_map, - normalize=self.normalize, + data_layout=self.data_layout, + preprocessor=self.preprocessor, ) logger.info(f"Aggregating inference logs") @@ -776,18 +777,19 @@ def init_data_loaders(self, cur_step: int) -> None: """ train_datasets = [ TorchTrainDataset( - src=src, + input_source=source, + label_source=None, prognostic_var_names=self.prognostic_var_names, boundary_var_names=self.boundary_var_names, hist=self.hist, steps=cur_step, normalize_before_mask=self.normalize_before_mask, - masked_fill_value=self.normalize_fill_value, + masked_fill_value=self.masked_fill_value, stride=stride, concurrent_compute_=self.concurrent_compute, ) for stride in self.data_stride - for src in self.data_container.train_sources + for source in self.data_bundle.train_sources ] # Validation is always evaluated on the primary source. This keeps the @@ -795,13 +797,14 @@ def init_data_loaders(self, cur_step: int) -> None: # regardless of the set of resolutions used for training. val_datasets = [ TorchTrainDataset( - src=self.data_container.val_sources[0], + input_source=self.data_bundle.val_sources[0], + label_source=None, prognostic_var_names=self.prognostic_var_names, boundary_var_names=self.boundary_var_names, hist=self.hist, steps=1, # current_step set to 1 for validation normalize_before_mask=self.normalize_before_mask, - masked_fill_value=self.normalize_fill_value, + masked_fill_value=self.masked_fill_value, stride=stride, concurrent_compute_=self.concurrent_compute, ) @@ -811,11 +814,11 @@ def init_data_loaders(self, cur_step: int) -> None: # Create datasets match self.loader_version: case TorchTrainDataset.FLAG: - train_data: torch.utils.data.Dataset[RawTrainData] = ConcatDataset( + host_train_dataset: torch.utils.data.Dataset[HostBatch] = ConcatDataset( train_datasets ) - val_data: torch.utils.data.Dataset[RawTrainData] = ConcatDataset( + host_val_dataset: torch.utils.data.Dataset[HostBatch] = ConcatDataset( val_datasets ) @@ -828,7 +831,7 @@ def init_data_loaders(self, cur_step: int) -> None: match self.loader_version: case TorchTrainDataset.FLAG: - collate_fn = collate_raw_train_data + collate_fn = collate_host_batches case _: raise NotImplementedError( f"Collate function not defined for loader version " @@ -838,7 +841,7 @@ def init_data_loaders(self, cur_step: int) -> None: # Create batch samplers - branch on distributed vs non-distributed # Group by resolution so batches stay homogeneous across configured sources. def group_key(ds): - return ds.prognostic_src.grid_size + return tuple(source.grid_size for source in ds.sources) if self.distributed is not None: # Distributed training @@ -891,8 +894,8 @@ def group_key(ds): # Create data loaders (same for both distributed and non-distributed) # When using batch_sampler, don't specify batch_size or sampler - train_dataloader = DataLoader( - train_data, + host_train_loader = DataLoader( + host_train_dataset, batch_sampler=train_batch_sampler, num_workers=self.num_workers, persistent_workers=self.persistent_workers and self.num_workers > 0, @@ -901,8 +904,8 @@ def group_key(ds): multiprocessing_context=self.mp_context, ) - val_dataloader = DataLoader( - val_data, + host_val_loader = DataLoader( + host_val_dataset, batch_sampler=val_batch_sampler, num_workers=self.num_workers, persistent_workers=self.persistent_workers and self.num_workers > 0, @@ -912,10 +915,8 @@ def group_key(ds): ) # Wrap dataloaders to handle GPU post-processing - self.train_loader = TrainDataLoader( - train_dataloader, train_datasets, self.device - ) - self.val_loader = TrainDataLoader(val_dataloader, val_datasets, self.device) + self.train_loader = BatchLoader(host_train_loader, train_datasets, self.device) + self.val_loader = BatchLoader(host_val_loader, val_datasets, self.device) def save_all_checkpoints(self, epoch: int, v_loss: float, inf_loss: float): with self._test_context(): diff --git a/src/samudra/utils/ctx.py b/src/samudra/utils/ctx.py index 6d9c2d7c..f7a2b5f6 100644 --- a/src/samudra/utils/ctx.py +++ b/src/samudra/utils/ctx.py @@ -11,7 +11,7 @@ @dataclasses.dataclass(frozen=True) -class GridContext: +class BatchGrid: """Grid-level context for model forward passes and loss computation. Bundles spatial metadata that travels alongside input tensors during training: diff --git a/src/samudra/utils/data.py b/src/samudra/utils/data.py index b8c26b49..9ac10106 100644 --- a/src/samudra/utils/data.py +++ b/src/samudra/utils/data.py @@ -4,17 +4,17 @@ import dataclasses import logging -import re from collections import defaultdict -from collections.abc import Callable +from collections.abc import Mapping, Sequence from functools import cached_property -from typing import TYPE_CHECKING, Literal, Self +from types import MappingProxyType +from typing import TYPE_CHECKING, Any, Literal, Protocol, Self, final import numpy as np import torch import xarray as xr from einops import rearrange -from jaxtyping import Bool, Float +from jaxtyping import Bool if TYPE_CHECKING: from samudra.config import TimeConfig @@ -22,7 +22,7 @@ from samudra.constants import ( BatchTimeSeriesOutput, BoundaryVarNames, - DatasetSpec, + DataLayout, DictSingleChannelVar, Grid, GridMask, @@ -35,39 +35,30 @@ PrognosticMask, PrognosticVarNames, SingleTimeSeriesOutput, - TensorMap, construct_metadata, ) from samudra.derived_variables import add_derived_variables -from samudra.utils.llc import _preferred_available_var, _var_without_level logger = logging.getLogger(__name__) +_CANONICAL_MASK_PREFIX = "mask_" +_STACKED_WET_MASK_VAR = "wetmask" -def _var_name_encode_level(var_name: str) -> bool: - """Check if the variable name encodes the level.""" - var_name_encodes_level = re.compile(r"_[0-9]+") - return bool(var_name_encodes_level.search(var_name)) - -def _is_compact(data: xr.Dataset, means: xr.Dataset, stds: xr.Dataset) -> bool: - return all( - not _var_name_encode_level(str(v)) - for d in [data, means, stds] - for v in d.keys() - ) - - -@dataclasses.dataclass +@dataclasses.dataclass(frozen=True) class Masks: - """A collection of masks to expose the ocean and mask land.""" + """Read-only mask metadata used to expose the ocean and mask land. + + Tensor contents are shared between canonical views for efficiency and must + be treated as immutable by callers. + """ prognostic: PrognosticMask boundary: GridMask def __post_init__(self): - self.prognostic = self.prognostic.bool() - self.boundary = self.boundary.bool() + object.__setattr__(self, "prognostic", self.prognostic.bool()) + object.__setattr__(self, "boundary", self.boundary.bool()) def prognostic_with_hist( self, hist: int @@ -75,125 +66,232 @@ def prognostic_with_hist( return torch.concat([self.prognostic] * (hist + 1), dim=0) -@dataclasses.dataclass -class DataSource: - """Data source for the model.""" +@dataclasses.dataclass(frozen=True) +class CanonicalReadRequest: + """A storage-independent request for canonical ocean-data planes. + + The shape of ``time_indices`` defines the leading dimensions of the returned + planes. Keeping this core request to NumPy makes it usable by Python and native + readers without importing xarray concepts into the boundary. + """ + + time_indices: np.ndarray + channels: tuple[str, ...] + + def __post_init__(self) -> None: + indices = np.asarray(self.time_indices) + if not np.issubdtype(indices.dtype, np.integer): + raise TypeError("Canonical time indices must be integers") + immutable_indices = indices.astype(np.int64, copy=True) + immutable_indices.setflags(write=False) + object.__setattr__(self, "time_indices", immutable_indices) + object.__setattr__(self, "channels", tuple(self.channels)) + + +@dataclasses.dataclass(frozen=True) +class ChannelStatistics: + """Normalization statistics aligned one-for-one with canonical channels.""" + + mean: np.ndarray + std: np.ndarray + + +class CanonicalReader(Protocol): + """Narrow storage seam implemented by xarray now and native readers later.""" + + @property + def channels(self) -> tuple[str, ...]: ... + + @property + def time(self) -> xr.DataArray: ... + + @property + def resolution(self) -> tuple[Lat, Lon]: ... + + def statistics(self, channels: tuple[str, ...]) -> ChannelStatistics: ... + + @property + def attrs(self) -> Mapping[str, Any]: ... + + def slice_time(self, time: "TimeConfig") -> Self: ... + + def read(self, request: CanonicalReadRequest) -> np.ndarray: ... + + def coordinates(self) -> Mapping[str, xr.DataArray]: ... + + def metadata(self, data_layout: DataLayout) -> dict: ... + + +@dataclasses.dataclass(frozen=True) +class _XarrayCanonicalReader: + """Private xarray implementation of the canonical read contract.""" - name: str data: xr.Dataset means: xr.Dataset stds: xr.Dataset - masks: Masks - dataset_spec: DatasetSpec + channels: tuple[str, ...] - @cached_property - def is_compact(self) -> bool: - """Check if the data source is compact.""" - return _is_compact(self.data, self.means, self.stds) + @property + def time(self) -> xr.DataArray: + return self.data.time.copy(deep=True) - @cached_property + @property def resolution(self) -> tuple[Lat, Lon]: return ( - torch.from_numpy(self.data.lat.values), - torch.from_numpy(self.data.lon.values), + torch.from_numpy(self.data.lat.values).clone(), + torch.from_numpy(self.data.lon.values).clone(), ) - @cached_property - def grid_size(self) -> GridSize: - res = self.resolution - return res[0].shape[0], res[1].shape[0] + def statistics(self, channels: tuple[str, ...]) -> ChannelStatistics: + self._validate_channels(channels) + return ChannelStatistics( + _flatten(self.means[list(channels)]), + _flatten(self.stds[list(channels)]), + ) - @cached_property - def spherical_area_weights(self) -> Grid: - return spherical_area_weights(self.data) + @property + def attrs(self) -> Mapping[str, Any]: + return MappingProxyType(dict(self.data.attrs)) + + def _validate_channels(self, channels: tuple[str, ...]) -> None: + missing = set(channels).difference(self.channels) + if missing: + raise KeyError(f"Canonical channels not found: {sorted(missing)}") + + def slice_time(self, time: "TimeConfig") -> Self: + return dataclasses.replace(self, data=self.data.sel(time=time.time_slice)) + + def read(self, request: CanonicalReadRequest) -> np.ndarray: + self._validate_channels(request.channels) + index_dims = [f"index_{i}" for i in range(request.time_indices.ndim)] + index = xr.DataArray(request.time_indices, dims=index_dims) + selected = self.data[list(request.channels)].isel(time=index) + + # Materialize one combined graph, rather than loading canonical channels + # independently. Compact level views share their base-array Dask keys, so + # the scheduler can read/decompress each physical chunk once per request. + values = ( + selected.to_array(dim="channel") + .transpose(*index_dims, "channel", "lat", "lon") + .to_numpy() + .astype(np.float32, copy=False) + ) + return values + + def coordinates(self) -> Mapping[str, xr.DataArray]: + return MappingProxyType( + { + str(name): coordinate.copy(deep=True) + for name, coordinate in self.data.coords.items() + } + ) - @cached_property - def metadata(self) -> dict: - return construct_metadata(self.data, self.dataset_spec) + def metadata(self, data_layout: DataLayout) -> dict: + return construct_metadata(self.data, data_layout) - def filter( - self, - var_names: PrognosticVarNames | BoundaryVarNames, - *, - prefix: str, - ) -> Self: - """Filter the data source to only include the specified variables (and levels). - If the dataset is compact, it will also filter the levels based on the - variable names (which encode the level in the name). +@final +@dataclasses.dataclass(frozen=True) +class CanonicalSource: + """A structurally immutable, read-capable view of canonical ocean data. - Args: - var_names: Variable names to filter. - prefix: Prefix for the new data source name. + Physical xarray layout is private to the reader. In particular, callers see + the same ordered channels for flat and compact OM4 stores. Channel selection + and time slicing return new views and never mutate the source. Tensor-valued + masks are shared, read-only metadata; mutating their contents is unsupported. + """ - Returns: - A new `DataSource` only with the filtered variables and levels. - """ - name = f"{prefix}[{self.name}]" - if self.is_compact: - parsed_var_names, levels = [], [] - for mangled_var_name in var_names: - if not _var_name_encode_level(mangled_var_name): - parsed_var_names.append(mangled_var_name) - continue - tokens = mangled_var_name.split("_") - var_name, level = tokens[0], int(tokens[1]) - - parsed_var_names.append(var_name) - # Build set of total levels - if level not in levels: - levels.append(level) - - data = self.data[parsed_var_names] - means = self.means[parsed_var_names] - stds = self.stds[parsed_var_names] - if levels: - data = data.isel(lev=levels) - means = means.isel(lev=levels) - stds = stds.isel(lev=levels) - - return dataclasses.replace( - self, name=name, data=data, means=means, stds=stds - ) + name: str + _reader: CanonicalReader + masks: Masks + data_layout: DataLayout - data = self.data[var_names] - means = self.means[var_names] - stds = self.stds[var_names] + @property + def reader(self) -> CanonicalReader: + """Return the storage reader so backends can decorate its read behavior.""" + return self._reader - return dataclasses.replace(self, name=name, data=data, means=means, stds=stds) + def with_reader(self, reader: CanonicalReader) -> Self: + """Return an equivalent source backed by a replacement reader.""" + if reader.channels != self.channels: + raise ValueError( + "Replacement reader channels must match the canonical source: " + f"expected {self.channels}, got {reader.channels}" + ) + return dataclasses.replace(self, _reader=reader) - def map( - self, - func: Callable[ - [xr.Dataset, xr.Dataset, xr.Dataset], - tuple[xr.Dataset, xr.Dataset, xr.Dataset], - ], - *, - suffix: str | None = None, + @classmethod + def from_canonical_datasets( + cls, + name: str, + data: xr.Dataset, + means: xr.Dataset, + stds: xr.Dataset, + masks: Masks, + data_layout: DataLayout, ) -> Self: - """Map the function over the data source.""" - if suffix is None: - suffix = func.__qualname__ - - data, means, stds = func(self.data.copy(), self.means.copy(), self.stds.copy()) + """Construct from datasets that are already in canonical channel form. - return dataclasses.replace( - self, name=f"{self.name}_{suffix}", data=data, means=means, stds=stds + Raw OM4 callers should use :meth:`from_datasets`. This factory remains + useful for focused in-memory tests and named preprocessing stages. + """ + channels = tuple(str(name) for name in means.data_vars) + if any("lev" in dataset.dims for dataset in (data, means, stds)): + raise ValueError("Canonical datasets cannot expose a 'lev' dimension") + if set(channels) - set(data.data_vars) or set(channels) - set(stds.data_vars): + raise ValueError("Canonical data, means, and stds have different channels") + return cls( + name=name, + _reader=_XarrayCanonicalReader( + data, + means[list(channels)], + stds[list(channels)], + channels, + ), + masks=masks, + data_layout=data_layout, ) - def map_data( - self, func: Callable[[xr.Dataset], xr.Dataset], *, suffix: str | None = None - ) -> Self: - """Map the function over just data in DataSource.""" - if suffix is None: - suffix = func.__qualname__ - data = func(self.data.copy()) - return dataclasses.replace(self, name=f"{self.name}_{suffix}", data=data) + @property + def channels(self) -> tuple[str, ...]: + return self._reader.channels - def slice(self, time: "TimeConfig") -> Self: + @property + def time(self) -> xr.DataArray: + return self._reader.time + + def statistics(self, channels: Sequence[str]) -> ChannelStatistics: + return self._reader.statistics(tuple(channels)) + + @property + def attrs(self) -> MappingProxyType[str, Any]: + return MappingProxyType(dict(self._reader.attrs)) + + @property + def resolution(self) -> tuple[Lat, Lon]: + # Readers return defensive coordinate tensors. Do not cache and expose a + # mutable tensor that could silently alter future callers' grid context. + return self._reader.resolution + + @cached_property + def grid_size(self) -> GridSize: + res = self.resolution + return res[0].shape[0], res[1].shape[0] + + @cached_property + def spherical_area_weights(self) -> Grid: + lat, lon = self.resolution + weights = torch.cos(torch.deg2rad(lat)).repeat(lon.shape[0], 1).t() + return weights / weights.sum() + + @cached_property + def metadata(self) -> dict: + return self._reader.metadata(self.data_layout) + + def slice_time(self, time: "TimeConfig") -> Self: """Slice the data source to only include the specified time slice.""" - data_time_min = self.data.time.values.min() - data_time_max = self.data.time.values.max() + data_time_min = self.time.values.min() + data_time_max = self.time.values.max() time_start = time.time_slice.start time_end = time.time_slice.stop if time_start > data_time_max or time_end < data_time_min: @@ -208,65 +306,26 @@ def slice(self, time: "TimeConfig") -> Self: f"{str(data_time_min)[:10]} to {str(data_time_max)[:10]}" ) - data = self.data.sel(time=time.time_slice) - return dataclasses.replace(self, name=f"{time=}[{self.name}]", data=data) - - # TODO(jder): delete this once we've de-duplicated InferenceDataset with TorchTrainDataset - def normalize(self, fill_nan=True, fill_value=0.0) -> xr.Dataset: - """Normalize input data.""" - norm = (self.data - self.means) / self.stds - if fill_nan: - norm = norm.fillna(fill_value) - return norm + return dataclasses.replace( + self, + name=f"{time=}[{self.name}]", + _reader=self._reader.slice_time(time), + ) - # TODO(jder): delete this once we've de-duplicated InferenceDataset with TorchTrainDataset - def normalize_with( - self, - data: torch.Tensor, - variable_axis: int = 0, - fill_nan=True, - fill_value=0.0, - ) -> torch.Tensor: - """Normalize input data treated as torch Tensors.""" - reshape_vars = [1] * data.ndim - reshape_vars[variable_axis] = -1 - - # TODO(alxmrs): Do we have to reshape twice? - if "lev" in self.means.dims: - means_np = ( - conditional_rearrange( - self.means, - "(variable lev)=var", - concat_dim="var", - ) - .rename({"var": "variable"}) - .to_numpy() - .reshape(-1) - ) - else: - means_np = self.means.to_array().to_numpy().reshape(-1) - if "lev" in self.stds.dims: - stds_np = ( - conditional_rearrange( - self.stds, - "(variable lev)=var", - concat_dim="var", - ) - .rename({"var": "variable"}) - .to_numpy() - .reshape(-1) - ) - else: - stds_np = self.stds.to_array().to_numpy().reshape(-1) + def read(self, time_indices: np.ndarray, channels: Sequence[str]) -> np.ndarray: + """Read canonical channels at integer time indices.""" + return self._reader.read(CanonicalReadRequest(time_indices, tuple(channels))) - means = torch.from_numpy(means_np).reshape(reshape_vars) - stds = torch.from_numpy(stds_np).reshape(reshape_vars) + def coordinates(self) -> dict[str, xr.DataArray]: + return dict(self._reader.coordinates()) - norm = (data - means) / stds - if fill_nan: - norm = norm.nan_to_num(nan=fill_value) - norm = norm.to(data.dtype) - return norm + def _xarray_datasets_for_testing( + self, + ) -> tuple[xr.Dataset, xr.Dataset, xr.Dataset]: + """Expose xarray fixtures without making them part of the public contract.""" + if not isinstance(self._reader, _XarrayCanonicalReader): + raise TypeError("This canonical dataset is not backed by xarray") + return self._reader.data, self._reader.means, self._reader.stds @classmethod def from_datasets( @@ -275,236 +334,74 @@ def from_datasets( means: xr.Dataset, stds: xr.Dataset, *, - dataset_spec: DatasetSpec, + data_layout: DataLayout, prognostic_var_names: PrognosticVarNames, boundary_var_names: BoundaryVarNames, - name: str = "DataSource", + name: str = "CanonicalSource", ) -> Self: - data, means, stds = validate_data( - data, - means, - stds, - dataset_spec=dataset_spec, - boundary_var_names=boundary_var_names, - ) + """Build a canonical reader from already-canonicalized xarray datasets.""" + channels = tuple(dict.fromkeys((*prognostic_var_names, *boundary_var_names))) + for dataset in (data, means, stds): + level_channels = [ + name + for name in channels + if name in dataset and "lev" in dataset[name].dims + ] + if level_channels: + raise ValueError( + "Canonical channels cannot expose a 'lev' dimension: " + f"{level_channels}" + ) masks = extract_wet_mask( data, prognostic_var_names, boundary_var_names, - dataset_spec=dataset_spec, ) + missing_data = set(channels).difference(data.data_vars) + missing_means = set(channels).difference(means.data_vars) + missing_stds = set(channels).difference(stds.data_vars) + if missing_data or missing_means or missing_stds: + raise ValueError( + "Canonical channels are missing: " + f"data={sorted(missing_data)}, means={sorted(missing_means)}, " + f"stds={sorted(missing_stds)}" + ) + return cls( name=name, - data=data, - means=means, - stds=stds, + _reader=_XarrayCanonicalReader( + data=data, + means=means[list(channels)], + stds=stds[list(channels)], + channels=channels, + ), masks=masks, - dataset_spec=dataset_spec, - ) - - -@dataclasses.dataclass -class OceanData: - """A slice of ocean data (boundary or prognostic) with normalization statistics. - - This dataclass bundles raw tensor data with the statistics needed to normalize it - and the mask needed to handle land/invalid values. It serves as an intermediary - representation used when constructing training `Example`s from raw xarray data. - - The typical workflow is: - 1. Load raw data from a `DataSource` via `from_data_source()` - 2. Slice to the desired time range with `with_time()` - 3. Apply normalization and masking with `normalize_and_mask()` - 4. Flatten time/variable dims to create the final `Input` or `Prognostic` tensor - - Attributes: - data: Raw ocean variable values with shape (batch, time, variable, lat, lon). - means: Per-variable means for normalization, shape (variable,). - stds: Per-variable standard deviations for normalization, shape (variable,). - mask: Boolean mask indicating valid ocean points (True) vs land (False), - broadcast-compatible with the variable dimension. - """ - - data: Float[torch.Tensor, "batch time variable lat lon"] - means: Float[torch.Tensor, " variable"] - stds: Float[torch.Tensor, " variable"] - mask: Bool[torch.Tensor, " variable"] - - @classmethod - def from_data_source( - cls, - data: Float[torch.Tensor, "batch time variable lat lon"], - mask: Float[torch.Tensor, " variable"], - src: DataSource, - ) -> Self: - means_torch = torch.from_numpy(_flatten(src.means)) - stds_torch = torch.from_numpy(_flatten(src.stds)) - return cls(data, means_torch, stds_torch, mask) - - def with_time(self, time_range: slice) -> Self: - """Slice data across the time dimension.""" - return dataclasses.replace(self, data=self.data[:, time_range, :, :, :]) - - def _normalize( - self, - data: Float[torch.Tensor, "batch time var lat lon"], - fill_nan: bool = True, - fill_value: float = 0.0, - ) -> Float[torch.Tensor, "batch time var lat lon"]: - """Normalize input data treated as torch Tensors.""" - norm = (data - self.means.view(1, 1, -1, 1, 1)) / self.stds.view(1, 1, -1, 1, 1) - if fill_nan: - norm = norm.nan_to_num(nan=fill_value) - norm = norm.to(data.dtype) - return norm - - def normalize_and_mask( - self, normalize_before_mask: bool, masked_fill_value: float - ) -> Float[torch.Tensor, "batch time var lat lon"]: - """Normalize and mask tensors.""" - tensor = self.data - if normalize_before_mask: - tensor = self._normalize(tensor) - tensor = torch.where(self.mask, tensor, masked_fill_value) - if not normalize_before_mask: - tensor = self._normalize(tensor) - return tensor - - def to(self, device: torch.device, non_blocking: bool = True) -> Self: - return dataclasses.replace( - self, - data=self.data.to(device, non_blocking=non_blocking), - means=self.means.to(device, non_blocking=non_blocking), - stds=self.stds.to(device, non_blocking=non_blocking), - mask=self.mask.to(device, non_blocking=non_blocking), + data_layout=data_layout, ) @dataclasses.dataclass -class DataSourceSplits: - train: DataSource - val: DataSource - inference: DataSource | None +class SourceSplits: + train: CanonicalSource + val: CanonicalSource + inference: CanonicalSource | None @dataclasses.dataclass -class DataContainer: - train_sources: list[DataSource] - val_sources: list[DataSource] - inference_source: DataSource | None +class DataBundle: + train_sources: list[CanonicalSource] + val_sources: list[CanonicalSource] + inference_source: CanonicalSource | None loader_version: LoaderVersion - dataset_spec: DatasetSpec - - # TODO: this is a bit of a footgun now that we have multiple kinds of sources - # and should be removed in favor of the appropriate source above. - @property - def primary_source(self) -> DataSource: - return self.train_sources[0] - - -def conditional_rearrange( - data: xr.Dataset, pattern: str, except_dim="lev", concat_dim="variable" -) -> xr.DataArray: - """Rearrange a Dataset using an einsum notation with and without a dimension. - - When a dataset has variables with a mixture of dimensions and an einsum-like - rearrange is applied on that dataset, it's common that the pattern will combinate - one too many variables. Sometimes, it's desirable to apply the rearrange pattern - on two versions of the data: one including variables with that dimension and one - without, and then concatenate them along a new dimension. - - For example, surface level boundary variables, which only occur at t0, should not be - combinatorially rearranged with depth variables that have multiple time steps. In - such a situation, this function can be used to apply a standard einsum rearrangement - to depth and surface variables, including and excluding variables who have a `time` - dimension, respectively. - - This method is stable: even if it creates a new number of dimensions, it will - preserve the order of the variables in the original dataset. - - Args: - data: The dataset to rearrange. - pattern: The einsum pattern to use for rearranging. - except_dim: The dimension to exclude from the pattern. - concat_dim: The dimension to concatenate along. - - Returns: - The combined, rearranged dataset as a `xarray.DataArray`. - """ - assert except_dim in pattern, f"{except_dim} must be in the pattern." - - all_vars = list(data.keys()) - - vars_with_dim = [v for v in data if except_dim in data[v].dims] - vars_without_dim = [v for v in data if except_dim not in data[v].dims] - - # Some of the `vars_without_dim` may need to appear before or behind `vars_with_dim` - # in the final data array. These lists help preserve the correct order of the vars, - # even after a rearrangement (i.e. merge to two or more dimensions). - back = [ - v - for v in vars_without_dim - if all_vars.index(v) > all_vars.index(vars_with_dim[0]) - ] - front = [ - v - for v in vars_without_dim - if all_vars.index(v) < all_vars.index(vars_with_dim[-1]) - ] - - data_with_dim = ( - data[vars_with_dim] - .to_array() - .einops.rearrange(pattern, dask="allowed") - .drop_vars(concat_dim, errors="ignore") - ) - data_without_dim = ( - data[vars_without_dim] - .to_array() - .einops.rearrange(pattern.replace(except_dim, ""), dask="allowed") - .drop_vars(concat_dim, errors="ignore") - ) - - da = xr.concat([data_with_dim, data_without_dim], dim=concat_dim) - - n_front = len(front) # e.g. n_front=2 - n_center = data_with_dim.sizes[concat_dim] # e.g. n_center=10 - n_back = len(back) # e.g. n_back=3 - - # In the `concat` above, we put all the `data_without_dim` vars at the end. Some of - # these need to be moved to the front, and the rest stays at the back. Here, we - # compute a list of indices that will sort the data in the correct order. - # - # e.g. with the example constants above, order would look like: - # array([10, 11, 0, 1, 2, 3, 4, 5, 6, 7, 8, 9, 12, 13, 14]) - order = np.concatenate( - ( - # Moves vars to the front - np.roll(np.arange((new_front := n_center + n_front)), n_front), - np.arange(new_front, new_front + n_back), # rest of vars - ) - ) - order_da = xr.DataArray(order, dims=concat_dim) - - return da.sortby(order_da) + data_layout: DataLayout def _flatten(ds: xr.Dataset) -> np.ndarray: - """Flatten a means/stds dataset into a 1-D array with one entry per channel. - - Variables with a `lev` dim are flattened with `lev` first; variables without - a `lev` dim are flattened as-is. This matches the channel layout of the - prognostic / boundary tensors, which `to_array().reshape(-1)` would otherwise - mis-align by broadcasting surface-only variables across all depth levels. - """ - flattened = [] - for name in ds.data_vars: - var = ds[name] - if "lev" in var.dims: - var = var.transpose("lev", ...) - flattened.append(var.to_numpy().reshape(-1)) - return np.concatenate(flattened) + """Flatten scalar statistics already aligned to canonical channel order.""" + if "lev" in ds.dims: + raise ValueError("Canonical statistics cannot expose a 'lev' dimension") + return ds.to_array().to_numpy().reshape(-1) def _level_index_from_var_name(var_name: str) -> int: @@ -512,11 +409,20 @@ def _level_index_from_var_name(var_name: str) -> int: return int(suffix) if suffix.isdigit() else 0 +def _var_without_level(var_name: str) -> str: + suffix = var_name.rsplit("_", maxsplit=1)[-1] + return var_name.rsplit("_", maxsplit=1)[0] if suffix.isdigit() else var_name + + +def _preferred_available_var( + data: xr.Dataset, candidates: tuple[str, ...] +) -> str | None: + return next((name for name in candidates if name in data.data_vars), None) + + def _mask_var_for_data_var( data: xr.Dataset, var_name: str, - *, - dataset_spec: DatasetSpec, ) -> str: level = _level_index_from_var_name(var_name) base_var = _var_without_level(var_name) @@ -526,19 +432,17 @@ def _mask_var_for_data_var( "V": (f"mask_s_{level}", f"hFacS_{level}"), "oceTAUY": ("mask_s_0", "hFacS_0"), } - if dataset_spec.type == "llc" and base_var in llc_staggered_masks: + if base_var in llc_staggered_masks: if mask_var := _preferred_available_var(data, llc_staggered_masks[base_var]): return mask_var - return dataset_spec.mask_vars[level] + return f"{_CANONICAL_MASK_PREFIX}{level}" def _mask_array_for_data_var( data: xr.Dataset, var_name: str, - *, - dataset_spec: DatasetSpec, ) -> np.ndarray: - mask = data[_mask_var_for_data_var(data, var_name, dataset_spec=dataset_spec)] + mask = data[_mask_var_for_data_var(data, var_name)] if "time" in mask.dims: mask = mask.isel(time=0) return mask.to_numpy() @@ -548,24 +452,17 @@ def extract_wet_mask( data: xr.Dataset, prognostic_var_names: PrognosticVarNames, boundary_var_names: BoundaryVarNames, - *, - dataset_spec: DatasetSpec, ) -> Masks: """A mask for where the oceans are. Water is wet.""" - data_ = flatten_masks(data, dataset_spec=dataset_spec) + data_ = flatten_masks(data) wet_inp_np = np.stack( - [ - _mask_array_for_data_var(data_, var_name, dataset_spec=dataset_spec) - for var_name in prognostic_var_names - ] + [_mask_array_for_data_var(data_, var_name) for var_name in prognostic_var_names] ) boundary_mask_vars = [ - _mask_var_for_data_var(data_, var_name, dataset_spec=dataset_spec) - for var_name in boundary_var_names + _mask_var_for_data_var(data_, var_name) for var_name in boundary_var_names ] boundary_masks = [ - _mask_array_for_data_var(data_, var_name, dataset_spec=dataset_spec) - for var_name in boundary_var_names + _mask_array_for_data_var(data_, var_name) for var_name in boundary_var_names ] wet_surface_mask_np = ( np.stack(boundary_masks) @@ -580,41 +477,38 @@ def extract_wet_mask( def flatten_masks( data: xr.Dataset, - dataset_spec: DatasetSpec, ) -> xr.Dataset: """Adds level-wise mask variables from the stacked wet mask.""" data_ = data.copy() - mask_vars = list(dataset_spec.mask_vars) - if mask_vars[0] not in data_.variables: - assert dataset_spec.mask_all_levels_var in data_.variables, ( + if f"{_CANONICAL_MASK_PREFIX}0" not in data_.variables: + assert _STACKED_WET_MASK_VAR in data_.variables, ( "Wet mask cannot be constructed without " "either the wetmask variable or the level-wise masks" ) - wet_mask = data_[dataset_spec.mask_all_levels_var] - for i, mask_var in enumerate(mask_vars): - data_[mask_var] = wet_mask.isel(lev=i) + wet_mask = data_[_STACKED_WET_MASK_VAR] + for i in range(wet_mask.sizes["lev"]): + data_[f"{_CANONICAL_MASK_PREFIX}{i}"] = wet_mask.isel(lev=i) - data_ = data_.drop_vars(dataset_spec.mask_all_levels_var) + data_ = data_.drop_vars(_STACKED_WET_MASK_VAR) return data_ def unflatten_masks( data: xr.Dataset, - dataset_spec: DatasetSpec, + num_levels: int, ) -> xr.Dataset: """Adds a stacked wet mask `xarray.DataArray` from level-wise mask variables.""" data_ = data.copy() - mask_vars = list(dataset_spec.mask_vars) - if dataset_spec.mask_all_levels_var not in data_.variables: + mask_vars = [f"{_CANONICAL_MASK_PREFIX}{i}" for i in range(num_levels)] + if _STACKED_WET_MASK_VAR not in data_.variables: assert mask_vars[0] in data_.variables, "Wet mask must have masks as data vars!" - wetmask = data_[mask_vars].to_array( - dim="lev", name=dataset_spec.mask_all_levels_var - ) + wetmask = data_[mask_vars].to_array(dim="lev", name=_STACKED_WET_MASK_VAR) - data_[dataset_spec.mask_all_levels_var] = wetmask.assign_coords(lev=data_.lev) + lev = data_.coords.get("lev", np.arange(len(mask_vars))) + data_[_STACKED_WET_MASK_VAR] = wetmask.assign_coords(lev=lev) data_ = data_.drop_vars(mask_vars) return data_ @@ -671,7 +565,7 @@ def spherical_area(data: xr.Dataset) -> Grid: return torch.from_numpy(areas) -def get_inference_steps(data_source: DataSource, hist: int = 1): +def get_inference_steps(data_source: CanonicalSource, hist: int = 1): """ Get the number of inference/rollout steps for the given time configuration. @@ -682,7 +576,7 @@ def get_inference_steps(data_source: DataSource, hist: int = 1): Returns: num_steps: Total number of rolled-out inferences which fit into the time range """ - num_steps = data_source.data.time.size + num_steps = data_source.time.size # Might have extra remaining days, so we remove them mod = num_steps % (hist + 1) @@ -693,21 +587,21 @@ def get_inference_steps(data_source: DataSource, hist: int = 1): def convert_tensor_out_to_dict( tensor_out: torch.Tensor, *, - tensor_map: TensorMap, + data_layout: DataLayout, ) -> DictSingleChannelVar: assert tensor_out.ndim == 5 - assert tensor_out.shape[2] == len(tensor_map.prognostic_var_names) + assert tensor_out.shape[2] == len(data_layout.prognostic_var_names) out_dict = {} - for i, var in enumerate(tensor_map.prognostic_var_names): + for i, var in enumerate(data_layout.prognostic_var_names): out_dict[var] = tensor_out[:, :, i] - out_dict.update(add_derived_variables(tensor_out, tensor_map=tensor_map)) + out_dict.update(add_derived_variables(tensor_out, data_layout=data_layout)) return out_dict def get_aggregator_dicts( data: Prognostic | Input, - normalize: "Normalize", - tensor_map: TensorMap, + preprocessor: "BatchPreprocessor", + data_layout: DataLayout, wet: torch.Tensor, long_rollout: bool, input_type: Literal["prognostic", "input"] = "prognostic", @@ -735,13 +629,13 @@ def get_aggregator_dicts( # Get normalized dict data_normalized = data_reshaped.clone() data_normalized = torch.where(wet == 0, float("nan"), data_normalized) - data_dict = convert_tensor_out_to_dict(data_normalized, tensor_map=tensor_map) + data_dict = convert_tensor_out_to_dict(data_normalized, data_layout=data_layout) # Unnormalize - data_unnorm = normalize.unnormalize_tensor_prognostic( + data_unnorm = preprocessor.unnormalize_tensor_prognostic( data_reshaped, fill_value=float("nan") ) # Get unnormalized dict - data_unnorm_dict = convert_tensor_out_to_dict(data_unnorm, tensor_map=tensor_map) + data_unnorm_dict = convert_tensor_out_to_dict(data_unnorm, data_layout=data_layout) return data_dict, data_unnorm_dict @@ -777,7 +671,7 @@ def compute_anomalies( def with_level_index_vars( data: xr.Dataset, - dataset_spec: DatasetSpec, + depth_levels: Sequence[float], ) -> xr.Dataset: """ Ensure variable names use a depth level index, not depth level value. @@ -792,7 +686,7 @@ def with_level_index_vars( var_split = var_str.split("_lev_") var = var_split[0] lev_in_depth = float(var_split[1].replace("_", ".")) - lev_in_depth_idx = dataset_spec.depth_levels.index(lev_in_depth) + lev_in_depth_idx = depth_levels.index(lev_in_depth) data_copy = data_copy.rename({var_str: f"{var}_{lev_in_depth_idx!s}"}) return data_copy @@ -800,23 +694,23 @@ def with_level_index_vars( def with_depth_value_vars( data: xr.Dataset, - dataset_spec: DatasetSpec, + data_layout: DataLayout, ) -> xr.Dataset: """Inverse of `with_level_index_vars`: name 3D variables by depth value. Renames the depth-resolved prognostic variables (``_``) back to the OM4 ``_lev_`` form (e.g. ``thetao_0`` -> ``thetao_lev_2_5``). Which variables to rename is read directly off - ``dataset_spec.prognostic_var_names`` rather than inferred from the data, so + ``data_layout.prognostic_var_names`` rather than inferred from the data, so per-level masks (``mask_``) and level-free prognostics (e.g. ``zos``) are never mistaken for depth-resolved variables by name alone. """ renames = {} - for var_name in dataset_spec.prognostic_var_names: + for var_name in data_layout.prognostic_var_names: base, _, idx = var_name.rpartition("_") if not (base and idx.isdigit()): continue # level-free prognostic variable, e.g. zos - depth = dataset_spec.depth_levels[int(idx)] + depth = data_layout.depth_levels[int(idx)] depth_str = str(depth).replace(".", "_") if var_name in data.variables: renames[var_name] = f"{base}_lev_{depth_str}" @@ -836,93 +730,118 @@ def with_lat_lon_coords(data: xr.Dataset) -> xr.Dataset: return data_copy -def validate_data( - data: xr.Dataset, - means: xr.Dataset, - stds: xr.Dataset, - dataset_spec: DatasetSpec, - boundary_var_names: BoundaryVarNames, -) -> tuple[xr.Dataset, xr.Dataset, xr.Dataset]: - """Validate the data such that we have the correct format for training.""" - is_compact = _is_compact(data, means, stds) +class BatchPreprocessor: + def __init__( + self, + source: CanonicalSource, + prognostic_var_names: Sequence[str], + boundary_var_names: Sequence[str], + *, + normalize_before_mask: bool = True, + masked_fill_value: float = 0.0, + ) -> None: + """Prepare canonical host tensors for models and restore physical values.""" + self.prognostic_mask = source.masks.prognostic + self.boundary_mask = source.masks.boundary + self.normalize_before_mask = normalize_before_mask + self.masked_fill_value = masked_fill_value + + # Pre-compute arrays for faster tensor normalization. Canonicalization has + # already aligned every scalar statistic to one ordered logical channel. + prognostic_statistics = source.statistics(prognostic_var_names) + boundary_statistics = source.statistics(boundary_var_names) + self._prognostic_mean_np = prognostic_statistics.mean + self._prognostic_std_np = prognostic_statistics.std + self._boundary_mean_np = boundary_statistics.mean + self._boundary_std_np = boundary_statistics.std + + @staticmethod + def _reshape_statistics(statistics: torch.Tensor, ndim: int) -> torch.Tensor: + shape = [1] * ndim + shape[-3] = -1 + return statistics.reshape(shape) - if is_compact: - data = with_lat_lon_coords(data) - else: - data = ( - data.pipe(flatten_masks, dataset_spec=dataset_spec) - .pipe(with_level_index_vars, dataset_spec=dataset_spec) - .pipe(with_lat_lon_coords) + @classmethod + def _normalize_tensor( + cls, + data: torch.Tensor, + mean: np.ndarray, + std: np.ndarray, + *, + fill_nan: bool = True, + fill_value: float = 0.0, + ) -> torch.Tensor: + if data.shape[-3] != mean.shape[0]: + raise ValueError( + f"Expected {mean.shape[0]} variable channels, got {data.shape[-3]}" + ) + tensor_mean = cls._reshape_statistics( + torch.from_numpy(mean).to(data.device, data.dtype), data.ndim ) + tensor_std = cls._reshape_statistics( + torch.from_numpy(std).to(data.device, data.dtype), data.ndim + ) + normalized = (data - tensor_mean) / tensor_std + if fill_nan: + normalized = normalized.nan_to_num(nan=fill_value) + return normalized.to(data.dtype) - # Check if data variables are in the right format - # This check is to ensure we convert data to the correct format - means = with_level_index_vars(means, dataset_spec=dataset_spec) - stds = with_level_index_vars(stds, dataset_spec=dataset_spec) - - # Check if any anomalies are needed to be computed - anomalies_vars = get_anomalies_vars(boundary_var_names) - out = ( - compute_anomalies(data, means, stds, anomalies_vars) - if anomalies_vars - else (data, means, stds) - ) - - return out + def _prepare( + self, + data: torch.Tensor, + *, + mean: np.ndarray, + std: np.ndarray, + mask: torch.Tensor, + device: torch.device, + ) -> Input: + tensor = data.to(device, non_blocking=True) + if tensor.ndim == 4: + tensor = tensor.unsqueeze(0) + elif tensor.ndim != 5: + raise ValueError(f"Expected 4D or 5D canonical planes, got {tensor.ndim}D") + mask = mask.to(device, non_blocking=True) + if self.normalize_before_mask: + tensor = self._normalize_tensor(tensor, mean, std) + tensor = torch.where(mask, tensor, self.masked_fill_value) + if not self.normalize_before_mask: + tensor = self._normalize_tensor(tensor, mean, std) + return rearrange( + tensor, "batch time variable lat lon -> batch (time variable) lat lon" + ) + def prepare_prognostic( + self, data: torch.Tensor, device: torch.device + ) -> Prognostic: + return self._prepare( + data, + mean=self._prognostic_mean_np, + std=self._prognostic_std_np, + mask=self.prognostic_mask, + device=device, + ) -class Normalize: - def __init__( - self, - src: DataSource, - prognostic_var_names: PrognosticVarNames, - boundary_var_names: BoundaryVarNames, - ) -> None: - """Store normalization parameters and pre-compute numpy arrays.""" - prognostic_src = src.filter(prognostic_var_names, prefix="prognostic") - boundary_src = src.filter(boundary_var_names, prefix="boundary") - self.prognostic_mean = prognostic_src.means - self.prognostic_std = prognostic_src.stds - self.boundary_mean = boundary_src.means - self.boundary_std = boundary_src.stds - self.wet_mask = src.masks.prognostic - self.wet_mask_surface = src.masks.boundary - - # Pre-compute numpy arrays for faster access. When a dataset mixes - # variables with and without a `lev` dim (e.g. thermo_dynamic_all has - # 4 depth-resolved variables plus surface-only `zos`), a plain - # `to_array().reshape(-1)` would broadcast the surface variable across - # all levels, producing too many channels. `_flatten` flattens each - # variable independently, matching the prognostic tensor channel layout. - self._prognostic_mean_np = _flatten(self.prognostic_mean) - self._prognostic_std_np = _flatten(self.prognostic_std) - self._boundary_mean_np = _flatten(self.boundary_mean) - self._boundary_std_np = _flatten(self.boundary_std) - self._wet_mask_np = self.wet_mask.numpy() + def prepare_boundary(self, data: torch.Tensor, device: torch.device) -> Input: + return self._prepare( + data, + mean=self._boundary_mean_np, + std=self._boundary_std_np, + mask=self.boundary_mask, + device=device, + ) def normalize_tensor_prognostic( self, data: torch.Tensor, fill_nan=True, fill_value=0.0 ) -> torch.Tensor: - """Normalize prognostic tensor.""" - tensor_mean = torch.from_numpy(self._prognostic_mean_np).to( - data.device, data.dtype - ) - tensor_std = torch.from_numpy(self._prognostic_std_np).to( - data.device, data.dtype + """Normalize a prognostic tensor without masking or flattening.""" + return self._normalize_tensor( + data, + self._prognostic_mean_np, + self._prognostic_std_np, + fill_nan=fill_nan, + fill_value=fill_value, ) - expand_var_dim = [1] * data.ndim - expand_var_dim[-3] = -1 - assert data.shape[-3] == self._prognostic_mean_np.shape[0] - tensor_mean = tensor_mean.reshape(expand_var_dim) - tensor_std = tensor_std.reshape(expand_var_dim) - - norm = (data - tensor_mean) / tensor_std - if fill_nan: - norm = norm.nan_to_num(nan=fill_value) - norm = norm.to(data.dtype) - return norm - def unnormalize_tensor_prognostic( self, data: torch.Tensor, fill_value=float("nan") ) -> torch.Tensor: @@ -941,7 +860,9 @@ def unnormalize_tensor_prognostic( tensor_std = tensor_std.reshape(expand_var_dim) unnorm = data * tensor_std + tensor_mean - unnorm = torch.where(self.wet_mask.to(data.device) == 0, fill_value, unnorm) + unnorm = torch.where( + self.prognostic_mask.to(data.device) == 0, fill_value, unnorm + ) unnorm = unnorm.to(data.dtype) return unnorm @@ -962,7 +883,7 @@ def unnormalize_tensor_boundary( unnorm = data * tensor_std + tensor_mean unnorm = torch.where( - self.wet_mask_surface.to(data.device) == 0, fill_value, unnorm + self.boundary_mask.to(data.device) == 0, fill_value, unnorm ) unnorm = unnorm.to(data.dtype) return unnorm @@ -970,7 +891,7 @@ def unnormalize_tensor_boundary( @dataclasses.dataclass class LoadStats: - """Captures stats about loading a single TrainData object.""" + """Captures stats about loading a single ModelBatch object.""" load_time_seconds: float @@ -1008,7 +929,7 @@ def _parse_level(x) -> float: def stack_levels( data: xr.Dataset, - dataset_spec: DatasetSpec, + data_layout: DataLayout, ) -> xr.Dataset: """Reassemble a flattened OM4 dataset into analysis-ready, depth-stacked form. @@ -1022,8 +943,8 @@ def stack_levels( Implemented as the inverse of `with_level_index_vars` followed by the existing `compact_dataset` (stacks the ``_lev_`` form) and `unflatten_masks`. """ - data = with_depth_value_vars(data, dataset_spec) + data = with_depth_value_vars(data, data_layout) data = compact_dataset(data) - if dataset_spec.mask_vars[0] in data.variables: - data = unflatten_masks(data, dataset_spec=dataset_spec) + if "mask_0" in data.variables: + data = unflatten_masks(data, num_levels=len(data_layout.depth_levels)) return data diff --git a/src/samudra/utils/llc.py b/src/samudra/utils/llc.py index a5b3d72a..99a2640e 100644 --- a/src/samudra/utils/llc.py +++ b/src/samudra/utils/llc.py @@ -4,7 +4,7 @@ import xarray as xr -from samudra.constants import DatasetSpec +from samudra.constants import DataLayout, build_llc_layout def _rename_llc_level_index_vars(ds: xr.Dataset) -> xr.Dataset: @@ -24,7 +24,7 @@ def _rename_llc_level_index_vars(ds: xr.Dataset) -> xr.Dataset: def _flatten_llc_level_vars( data: xr.Dataset, *, - dataset_spec: DatasetSpec, + num_levels: int, ) -> xr.Dataset: """Flatten LLC level dimensions into level-indexed variables. @@ -39,15 +39,13 @@ def _flatten_llc_level_vars( continue n_levels = data_copy[name].sizes["lev"] - expected_levels = len(dataset_spec.depth_i_levels) - if n_levels != expected_levels: + if n_levels != num_levels: raise ValueError( - f"Expected {expected_levels} levels for LLC variable {name}, got " - f"{n_levels}" + f"Expected {num_levels} levels for LLC variable {name}, got {n_levels}" ) - for index, lev in enumerate(dataset_spec.depth_i_levels): - data_copy[f"{name}_{lev}"] = data_copy[name].isel(lev=index) + for index in range(num_levels): + data_copy[f"{name}_{index}"] = data_copy[name].isel(lev=index) data_copy = data_copy.drop_vars(name) return data_copy @@ -102,25 +100,25 @@ def canonicalize_llc_datasets( i_end: int, j_start: int, j_end: int, - dataset_spec: DatasetSpec, -) -> tuple[xr.Dataset, xr.Dataset, xr.Dataset]: + prognostic_vars_key: str, + boundary_vars_key: str, +) -> tuple[xr.Dataset, xr.Dataset, xr.Dataset, DataLayout]: """Standardize raw LLC inputs to the common non-compact loader layout. For example, a selected ``Theta(time, face, k, j, i)`` crop becomes level-indexed ``Theta_0(time, y, x)``, ``Theta_50(time, y, x)``, and the remaining configured levels; statistics names undergo the same renaming. """ + data_layout = build_llc_layout(prognostic_vars_key, boundary_vars_key) data_copy = data.copy() requested_data_vars = { _var_without_level(var_name) for var_name in ( - dataset_spec.prognostic_var_names + dataset_spec.boundary_var_names + data_layout.prognostic_var_names + data_layout.boundary_var_names ) } - requested_data_vars.update( - [dataset_spec.mask_all_levels_var, "mask_c", *dataset_spec.mask_vars] - ) + requested_data_vars.update(["wetmask", "mask_c"]) requested_data_vars.update(_llc_staggered_mask_vars(data_copy, requested_data_vars)) data_copy = data_copy[ [name for name in data_copy.data_vars if name in requested_data_vars] @@ -179,8 +177,8 @@ def canonicalize_llc_datasets( for old, new in { "k": "lev", "mask_c": "wetmask", - "i": "x", - "j": "y", + "i": "lon", + "j": "lat", }.items() if old in data_copy.dims or old in data_copy.variables @@ -191,6 +189,15 @@ def canonicalize_llc_datasets( means_copy = _rename_llc_level_index_vars(means.copy()) stds_copy = _rename_llc_level_index_vars(stds.copy()) - data_copy = _flatten_llc_level_vars(data_copy, dataset_spec=dataset_spec) + data_copy = _flatten_llc_level_vars( + data_copy, num_levels=len(data_layout.depth_levels) + ) + mask_rename_map = { + f"wetmask_{level}": f"mask_{level}" + for level in range(len(data_layout.depth_levels)) + if f"wetmask_{level}" in data_copy + } + if mask_rename_map: + data_copy = data_copy.rename(mask_rename_map) - return data_copy, means_copy, stds_copy + return data_copy, means_copy, stds_copy, data_layout diff --git a/src/samudra/utils/logging.py b/src/samudra/utils/logging.py index 27b214ae..02f8e90b 100644 --- a/src/samudra/utils/logging.py +++ b/src/samudra/utils/logging.py @@ -24,9 +24,9 @@ logger = logging.getLogger(__name__) if TYPE_CHECKING: - from samudra.datasets import TrainData, TrainDataLoader + from samudra.datasets import BatchLoader, ModelBatch from samudra.models.base import BaseModel - from samudra.utils.ctx import GridContext + from samudra.utils.ctx import BatchGrid def handle_logging(debug: bool, output_dir: Path): @@ -161,7 +161,7 @@ def add_meter(self, name, meter): def log_every( self, - data_loader: "TrainDataLoader", + data_loader: "BatchLoader", print_freq, header=None, ): @@ -230,7 +230,7 @@ class _ForwardOnceWrapper(torch.nn.Module): def __init__( self, model: "BaseModel | DistributedDataParallel", - ctx: "GridContext", + ctx: "BatchGrid", ) -> None: super().__init__() self._underlying: BaseModel = getattr(model, "module", model) # type: ignore @@ -242,19 +242,21 @@ def forward(self, prognostic: Prognostic, boundary: Boundary) -> torch.Tensor: def get_model_summary( - model: "BaseModel | DistributedDataParallel", data: "TrainData | None", debug: bool + model: "BaseModel | DistributedDataParallel", + batch: "ModelBatch | None", + debug: bool, ) -> None: model_parameters = filter(lambda p: p.requires_grad, model.parameters()) params = sum([np.prod(p.size()) for p in model_parameters]) logger.info(f"Number of parameters: {params}") depth = 10 if debug else 2 # we pass verbose = 0 because we log the summary ourselves - if data is not None: - # TrainData is a complex wrapper that torchinfo cannot traverse. + if batch is not None: + # ModelBatch is a complex wrapper that torchinfo cannot traverse. # Extract the initial prognostic + boundary and wrap the model to # use forward_once. - prog_tensor, boundary_tensor = data.get_initial_input() - wrapper = _ForwardOnceWrapper(model, data.ctx) + prog_tensor, boundary_tensor = batch.get_initial_input() + wrapper = _ForwardOnceWrapper(model, batch.ctx) logger.info( summary( wrapper, diff --git a/src/samudra/utils/loss.py b/src/samudra/utils/loss.py index e2c9c562..26da6fd9 100644 --- a/src/samudra/utils/loss.py +++ b/src/samudra/utils/loss.py @@ -9,7 +9,7 @@ import torch.nn.functional as F from jaxtyping import Float -from samudra.utils.ctx import GridContext +from samudra.utils.ctx import BatchGrid LossFn = Callable[[torch.Tensor, torch.Tensor], torch.Tensor] @@ -19,7 +19,7 @@ def __call__( self, pred: torch.Tensor, target: torch.Tensor, - ctx: GridContext, + ctx: BatchGrid, ) -> torch.Tensor: ... @@ -47,7 +47,7 @@ def loss_fn_from_metric(metric: LossMetric) -> LossFnWithContext: def loss_fn_with_ctx( pred: torch.Tensor, target: torch.Tensor, - ctx: GridContext, + ctx: BatchGrid, ) -> torch.Tensor: wet = ctx.label_mask.to(device=pred.device) pred = pred * wet @@ -176,7 +176,7 @@ def __call__( self, pred: Float[torch.Tensor, "batch hist*var lat lon"], target: Float[torch.Tensor, "batch hist*var lat lon"], - ctx: GridContext, + ctx: BatchGrid, ) -> Float[torch.Tensor, " hist*var"]: loss_with_history_channels: Float[torch.Tensor, " hist*var"] = self.loss_fn( pred, target, ctx @@ -255,7 +255,7 @@ def __call__( self, pred: Float[torch.Tensor, "batch hist*var lat lon"], target: Float[torch.Tensor, "batch hist*var lat lon"], - ctx: GridContext, + ctx: BatchGrid, ) -> Float[torch.Tensor, " hist*var"]: base_loss = self.loss_fn(pred, target, ctx) # Ensure mask is on the same device as pred for gradient computation diff --git a/src/samudra/utils/output.py b/src/samudra/utils/output.py index 8552faa0..c3653fa8 100644 --- a/src/samudra/utils/output.py +++ b/src/samudra/utils/output.py @@ -5,7 +5,7 @@ import torch import xarray as xr -from samudra.utils.ctx import GridContext +from samudra.utils.ctx import BatchGrid class TrainBatchOutput: @@ -22,7 +22,7 @@ def __init__( input_data: torch.Tensor, target_data: torch.Tensor, gen_data: torch.Tensor, - ctx: GridContext, + ctx: BatchGrid, ): super().__init__(loss, loss_per_channel) assert target_data.shape == gen_data.shape diff --git a/src/samudra/utils/samplers.py b/src/samudra/utils/samplers.py index 16fe7ce3..a95d7bda 100644 --- a/src/samudra/utils/samplers.py +++ b/src/samudra/utils/samplers.py @@ -127,17 +127,17 @@ def from_datasets( seed: Random seed for deterministic shuffling Examples: - - lambda ds: (ds._input_src.data.sizes['lat'], ds._input_src.data.sizes['lon']) # group by resolution - - lambda ds: ds._input_src.data.sizes['lat'] # group by latitude size only + - lambda ds: (ds._input_source.data.sizes['lat'], ds._input_source.data.sizes['lon']) # group by resolution + - lambda ds: ds._input_source.data.sizes['lat'] # group by latitude size only Returns: EquivalenceGroupBatchSampler configured to group by the provided key - Example: + RolloutStep: >>> # Group datasets by resolution, allowing different strides to be batched together >>> sampler = EquivalenceGroupBatchSampler.from_datasets( ... datasets=dataset_list, - ... group_key=lambda ds: ds.prognostic_src.grid_size, + ... group_key=lambda ds: tuple(source.grid_size for source in ds.sources), ... batch_size=32, ... shuffle=True, ... drop_last=True, diff --git a/src/samudra/utils/train.py b/src/samudra/utils/train.py index 37c5aff4..390a1426 100644 --- a/src/samudra/utils/train.py +++ b/src/samudra/utils/train.py @@ -9,7 +9,7 @@ import torch from xarray_einstats.einops import rearrange # noqa: F401 -from samudra.datasets import InferenceDataset, RawTrainData +from samudra.datasets import HostBatch, InferenceDataset from samudra.utils.data import LoadStats @@ -20,18 +20,18 @@ def pairwise(iterable): return zip(a, b) -def collate_raw_train_data(data: Sequence[RawTrainData]) -> RawTrainData: - batched_data = RawTrainData(data[0].dataset_id) +def collate_host_batches(data: Sequence[HostBatch]) -> HostBatch: + batched_data = HostBatch(data[0].dataset_id) assert all(d.dataset_id == batched_data.dataset_id for d in data), ( "we don't support heterogenous batches yet" ) - steps = len(data[0].raw_data) + steps = len(data[0].steps) for step in range(steps): - input_ = torch.stack([d.raw_data[step][0] for d in data]) - boundary = torch.stack([d.raw_data[step][1] for d in data]) - label = torch.stack([d.raw_data[step][2] for d in data]) - batched_data.insert(input_, boundary, label) + input_ = torch.stack([d.steps[step][0] for d in data]) + boundary = torch.stack([d.steps[step][1] for d in data]) + label = torch.stack([d.steps[step][2] for d in data]) + batched_data.append(input_, boundary, label) stats = LoadStats.accumulated( [d.load_stats for d in data if d.load_stats is not None] diff --git a/src/samudra/utils/train_progress.py b/src/samudra/utils/train_progress.py index 1b16e16b..c53a5c17 100644 --- a/src/samudra/utils/train_progress.py +++ b/src/samudra/utils/train_progress.py @@ -10,7 +10,7 @@ import torch -from samudra.datasets import TrainData +from samudra.datasets import ModelBatch @dataclasses.dataclass @@ -51,19 +51,21 @@ class TrainBatchProgress: batch_seconds: float = 0.0 @classmethod - def from_train_data(cls, data: TrainData, world_size: int) -> "TrainBatchProgress": - label = data.get_label(0) + def from_model_batch( + cls, batch: ModelBatch, world_size: int + ) -> "TrainBatchProgress": + label = batch.get_label(0) ( local_batch_size, output_channels, output_grid_lat, output_grid_lon, ) = label.shape - input_grid_lat = data.ctx.input_resolution_cpu[0].shape[0] - input_grid_lon = data.ctx.input_resolution_cpu[1].shape[0] + input_grid_lat = batch.ctx.input_resolution_cpu[0].shape[0] + input_grid_lon = batch.ctx.input_resolution_cpu[1].shape[0] sample_windows = local_batch_size * world_size - model_examples = sample_windows * len(data) + model_examples = sample_windows * len(batch) output_grid_cells = model_examples * output_grid_lat * output_grid_lon target_values = output_grid_cells * output_channels @@ -117,17 +119,17 @@ class TrainProgress: @contextlib.contextmanager def batch( - self, data: TrainData, *, world_size: int, device: torch.device + self, batch: ModelBatch, *, world_size: int, device: torch.device ) -> Iterator[TrainBatchProgress]: - batch = TrainBatchProgress.from_train_data(data, world_size) + batch_progress = TrainBatchProgress.from_model_batch(batch, world_size) _synchronize_cuda_if_needed(device) batch_start_time = time.perf_counter() - yield batch + yield batch_progress _synchronize_cuda_if_needed(device) - batch.batch_seconds = time.perf_counter() - batch_start_time - self.update(batch, world_size=world_size) + batch_progress.batch_seconds = time.perf_counter() - batch_start_time + self.update(batch_progress, world_size=world_size) def update(self, batch: TrainBatchProgress, *, world_size: int) -> None: self.sample_windows_seen += batch.sample_windows diff --git a/src/samudra/utils/wandb.py b/src/samudra/utils/wandb.py index 12a4b5fd..d81f772d 100644 --- a/src/samudra/utils/wandb.py +++ b/src/samudra/utils/wandb.py @@ -27,7 +27,7 @@ if TYPE_CHECKING: from samudra.config import AnyTopLevelConfig - from samudra.utils.data import DataContainer + from samudra.utils.data import DataBundle class WandBLogger(Multiton): @@ -40,10 +40,10 @@ def _initialize(self): def enabled(self): return self._enabled - def _make_config(self, cfg: "AnyTopLevelConfig", data_container: "DataContainer"): + def _make_config(self, cfg: "AnyTopLevelConfig", data_bundle: "DataBundle"): config = { - f"data_{i}/attrs": src.data.attrs - for i, src in enumerate(data_container.train_sources) + f"data_{i}/attrs": dict(source.attrs) + for i, source in enumerate(data_bundle.train_sources) } config.update(config=cfg.model_dump()) provenance_env = { @@ -68,7 +68,7 @@ def setup_run( self, checkpoint_path: str | None, cfg: "AnyTopLevelConfig", - data_container: "DataContainer", + data_bundle: "DataBundle", finetune: bool = False, ): """Set up a wandb run, either resuming from checkpoint or creating new run. @@ -76,17 +76,17 @@ def setup_run( Args: checkpoint_path: Path to checkpoint file, if resuming cfg: Configuration object - data_container: Data container to log attributes of + data_bundle: Data container to log attributes of finetune: Whether this is a finetuning run Returns: tuple: (wandb_id, wandb_name) """ if not checkpoint_path: - return self._init_new_run(cfg, data_container) + return self._init_new_run(cfg, data_bundle) if finetune: - return self._init_new_run(cfg, data_container) + return self._init_new_run(cfg, data_bundle) if not self._enabled: return None, None @@ -102,7 +102,7 @@ def setup_run( try: self.init( - config=self._make_config(cfg, data_container), + config=self._make_config(cfg, data_bundle), name=wandb_name, dir=cfg.experiment.output_dir, resume="must", @@ -112,7 +112,7 @@ def setup_run( except Exception: # If resume fails, start new run self.init( - config=self._make_config(cfg, data_container), + config=self._make_config(cfg, data_bundle), name=wandb_name, dir=cfg.experiment.output_dir, **cfg.experiment.wandb.model_dump(), @@ -120,19 +120,19 @@ def setup_run( return wandb_id, wandb_name - def _init_new_run(self, cfg: "AnyTopLevelConfig", data_container: "DataContainer"): + def _init_new_run(self, cfg: "AnyTopLevelConfig", data_bundle: "DataBundle"): """Initialize a new wandb run. Args: cfg: Configuration object - data_container: Data container to log attributes of + data_bundle: Data container to log attributes of Returns: tuple: (None, generated_name) for new run """ wandb_name = cfg.experiment.name if self._enabled: self.init( - config=self._make_config(cfg, data_container), + config=self._make_config(cfg, data_bundle), name=wandb_name, dir=cfg.experiment.output_dir, **cfg.experiment.wandb.model_dump(), diff --git a/src/samudra/utils/writer.py b/src/samudra/utils/writer.py index f2a7dac0..1367a0da 100644 --- a/src/samudra/utils/writer.py +++ b/src/samudra/utils/writer.py @@ -9,8 +9,8 @@ import xarray as xr from einops import rearrange -from samudra.constants import TensorMap -from samudra.utils.data import Normalize, stack_levels +from samudra.constants import DataLayout +from samudra.utils.data import BatchPreprocessor, stack_levels from samudra.utils.output import ModelInferenceOutput @@ -24,8 +24,8 @@ def __init__( hist: int, model_path: str | os.PathLike, time_chunk_size: int, - normalize: Normalize, - tensor_map: TensorMap, + preprocessor: BatchPreprocessor, + data_layout: DataLayout, ): self.pred_path = os.path.join(output_dir, "predictions.zarr") @@ -41,16 +41,16 @@ def __init__( self.model_path = model_path self.time_chunk_size = time_chunk_size - self.normalize = normalize - self.tensor_map = tensor_map + self.preprocessor = preprocessor + self.data_layout = data_layout - def record_batch(self, IO: ModelInferenceOutput): - pred_tensor = IO.prediction - pred_time = IO.time + def record_batch(self, inference_output: ModelInferenceOutput): + pred_tensor = inference_output.prediction + pred_time = inference_output.time pred_tensor = rearrange( pred_tensor, "n (hi c) h w -> (n hi) c h w", hi=self.hist + 1 ) - pred_tensor = self.normalize.unnormalize_tensor_prognostic( + pred_tensor = self.preprocessor.unnormalize_tensor_prognostic( pred_tensor, fill_value=0.0 ) if self.buffer is None: @@ -81,11 +81,11 @@ def write(self): per_level = xr.Dataset( { name: (["time", "y", "x"], buffer[:, channel, :, :]) - for channel, name in enumerate(self.tensor_map.prognostic_var_names) + for channel, name in enumerate(self.data_layout.prognostic_var_names) }, coords=coords, ) - ds = stack_levels(per_level, self.tensor_map.dataset_spec) + ds = stack_levels(per_level, self.data_layout) ds = ds.transpose("time", "lev", "y", "x", ...) ds.attrs["model_path"] = str(self.model_path) ds = ds.chunk({"time": self.time_chunk_size}) @@ -118,8 +118,7 @@ def _output_coords(self) -> dict: x_vals = np.asarray(src["lon"].values) # 1-D longitude axis ny, nx = y_vals.size, x_vals.size - spec = self.tensor_map.dataset_spec - n_levels = spec.num_prognostic_depth_levels + n_levels = self.data_layout.num_prognostic_depth_levels # 2-D lat/lon on the (y, x) grid, matching the ground-truth layout. Prefer # the real coordinates the source preserved (`lat_2d`/`lon_2d`), which are @@ -130,12 +129,13 @@ def _output_coords(self) -> dict: if "lat_2d" in src and "lon_2d" in src: lat2d = np.asarray(src["lat_2d"].values) lon2d = np.asarray(src["lon_2d"].values) - elif spec.grid_type == "gaussian": + elif self.data_layout.grid_type == "gaussian": lat2d = np.broadcast_to(y_vals[:, None], (ny, nx)).copy() lon2d = np.broadcast_to(x_vals[None, :], (ny, nx)).copy() else: raise ValueError( - f"Cannot build 2-D lat/lon for grid_type={spec.grid_type!r}: the " + "Cannot build 2-D lat/lon for " + f"grid_type={self.data_layout.grid_type!r}: the " "source coords carry no real 'lat_2d'/'lon_2d', and broadcasting the " "1-D axes is only valid on a 'gaussian' (rectilinear) grid. Preserve " "the true 2-D coordinates through preprocessing for curvilinear grids." @@ -146,7 +146,7 @@ def _output_coords(self) -> dict: "x": ("x", x_vals), "lat": (("y", "x"), lat2d), "lon": (("y", "x"), lon2d), - "lev": ("lev", np.array(spec.depth_levels[:n_levels])), + "lev": ("lev", np.array(self.data_layout.depth_levels[:n_levels])), } rename = {"lat": "y", "lon": "x"} diff --git a/src/samudra/viz/core.py b/src/samudra/viz/core.py index 0f778ba1..bf50dcf6 100644 --- a/src/samudra/viz/core.py +++ b/src/samudra/viz/core.py @@ -31,7 +31,7 @@ from matplotlib.ticker import FixedLocator, MaxNLocator, ScalarFormatter from tqdm.auto import tqdm -from samudra.constants import DatasetSpec, build_om4_spec +from samudra.constants import DataLayout, build_om4_layout from samudra.utils.data import ( spherical_area, spherical_area_weights, @@ -125,9 +125,9 @@ def __init__( } key1 = runs[0].name - # TODO: Support non-OM4 dataset specs in visualization. - self.dataset_spec = build_om4_spec() - levels = len(self.dataset_spec.depth_levels) + # TODO: Support non-OM4 data layouts in visualization. + self.data_layout = build_om4_layout() + levels = len(self.data_layout.depth_levels) groundtruth_rollout = groundtruth_rollout.sel(time=time_range) @@ -150,7 +150,7 @@ def __init__( # This function processes the ds_groundtruth and predictions for plotting # The predictions are loaded into pred_dict data, pred_dict = process_data( - groundtruth_rollout, pred_dict, dataset_spec=self.dataset_spec + groundtruth_rollout, pred_dict, data_layout=self.data_layout ) last_index = len(data.time) - 1 @@ -3749,7 +3749,7 @@ def isnan(x: xr.DataArray) -> xr.DataArray: def _combine_variables_by_level( - ds: xr.Dataset, combine_vars: list[str], dataset_spec: DatasetSpec + ds: xr.Dataset, combine_vars: list[str], data_layout: DataLayout ) -> xr.Dataset: """ Combine variables in the dataset along a new 'lev' dimension based on their suffix. @@ -3763,12 +3763,12 @@ def _combine_variables_by_level( xarray.Dataset: The dataset with combined variables and a new 'lev' dimension. """ for v in combine_vars: - level_numbers = [i for i in range(len(dataset_spec.depth_levels))] + level_numbers = [i for i in range(len(data_layout.depth_levels))] sorted_vars = [v + "_" + str(lev) for lev in level_numbers] if sorted_vars[0] not in ds.data_vars: continue combined = xr.concat([ds[var] for var in sorted_vars], dim="lev") - combined = combined.assign_coords(lev=list(dataset_spec.depth_levels)) + combined = combined.assign_coords(lev=list(data_layout.depth_levels)) ds[v] = combined ds = ds.drop_vars(sorted_vars) return ds @@ -3777,7 +3777,7 @@ def _combine_variables_by_level( def combine_variables_by_level( ds_groundtruth: xr.Dataset, pred_dict: dict[str, dict[str, Any]], - dataset_spec: DatasetSpec, + data_layout: DataLayout, ) -> tuple[xr.Dataset, dict[str, dict[str, Any]]]: """ Combine variables by level for ground truth and predictions. @@ -3790,11 +3790,11 @@ def combine_variables_by_level( xarray.Dataset, dict: Updated ground truth and prediction datasets. """ ds_groundtruth = _combine_variables_by_level( - ds_groundtruth, ["thetao", "so", "uo", "vo", "mask"], dataset_spec + ds_groundtruth, ["thetao", "so", "uo", "vo", "mask"], data_layout ) for key in pred_dict.keys(): pred_dict[key]["ds_prediction"] = _combine_variables_by_level( - pred_dict[key]["ds_prediction"], pred_dict[key]["ls"], dataset_spec + pred_dict[key]["ds_prediction"], pred_dict[key]["ls"], data_layout ) return ds_groundtruth, pred_dict @@ -3901,12 +3901,12 @@ def postprocess_for_plot( def process_data( data: xr.Dataset, pred_dict: dict[str, dict[str, Any]], - dataset_spec: DatasetSpec, + data_layout: DataLayout, ) -> tuple[xr.Dataset, dict[str, dict[str, Any]]]: """ Get plot ready OM4 data. """ - ds_groundtruth = with_level_index_vars(data, dataset_spec=dataset_spec) + ds_groundtruth = with_level_index_vars(data, depth_levels=data_layout.depth_levels) # Store ds_prediction copy_dict = deepcopy(pred_dict) @@ -3929,14 +3929,14 @@ def process_data( ### Combine Variables by level ds_groundtruth, pred_dict = combine_variables_by_level( - ds_groundtruth, pred_dict, dataset_spec + ds_groundtruth, pred_dict, data_layout ) ### Postprocess predictions for plotting ds_groundtruth, pred_dict = postprocess_for_plot( ds_groundtruth, ds_groundtruth.areacello, - np.array(dataset_spec.depth_thickness), + np.array(data_layout.depth_thickness), pred_dict, ) diff --git a/tests/conftest.py b/tests/conftest.py index bf4e232f..9ffbc3d4 100644 --- a/tests/conftest.py +++ b/tests/conftest.py @@ -16,25 +16,39 @@ from numpy.typing import ArrayLike, NDArray import samudra.constants as c -from samudra.config import TrainBackendConfig, TrainConfig +from samudra.config import Om4DataSourceConfig, TrainBackendConfig, TrainConfig from samudra.train import Trainer -from samudra.utils.data import DataSource, Masks, _is_compact, compact_dataset +from samudra.utils.data import CanonicalSource, Masks, compact_dataset from samudra.utils.multiton import MultitonScope REMOTE_DATA = "https://nyu1.osn.mghpcc.org/m2lines-pubs/Samudra/" DEFAULT_CONFIG = "test/train_default.yaml" SAMUDRA_MULTI_CONFIG = "test/train_samudra_multi.yaml" ALL_CONFIGS = [DEFAULT_CONFIG, "test/train_default_2step.yaml", SAMUDRA_MULTI_CONFIG] -TEST_DATASET_SPEC = c.build_om4_spec( +TEST_DATA_LAYOUT = c.build_om4_layout( prognostic_vars_key="thetao_1", boundary_vars_key="hfds", ) -TEST_FULL_DATASET_SPEC = c.build_om4_spec( +TEST_FULL_DATA_LAYOUT = c.build_om4_layout( prognostic_vars_key="thermo_dynamic_all", boundary_vars_key="tau_hfds_hfds_anom", ) +def canonicalize_mock_om4( + data: xr.Dataset, means: xr.Dataset, stds: xr.Dataset +) -> tuple[xr.Dataset, xr.Dataset, xr.Dataset]: + """Run test OM4 arrays through the production xarray canonicalizer.""" + canonicalizer = Om4DataSourceConfig.model_construct( + prognostic_vars_key="thermo_dynamic_all", + boundary_vars_key="hfds", + ) + canonical_data, canonical_means, canonical_stds, _ = ( + canonicalizer.canonicalize_datasets(data, means, stds) + ) + return canonical_data, canonical_means, canonical_stds + + def _om4_canonical_var_names(var_names: Iterable[str]) -> list[str]: out = [] for var_name in var_names: @@ -44,7 +58,7 @@ def _om4_canonical_var_names(var_names: Iterable[str]) -> list[str]: base_name, lev_depth = var_name.split("_lev_", maxsplit=1) depth_level = float(lev_depth.replace("_", ".")) - out.append(f"{base_name}_{TEST_DATASET_SPEC.depth_levels.index(depth_level)}") + out.append(f"{base_name}_{TEST_DATA_LAYOUT.depth_levels.index(depth_level)}") return out @@ -205,7 +219,7 @@ def to_coords(self) -> dict[str, xr.DataArray]: coords = { "lon": xr.DataArray(self.lng, dims=["lon"]), "lat": xr.DataArray(self.lat, dims=["lat"]), - "lev": xr.DataArray(np.array(TEST_DATASET_SPEC.depth_levels), dims=["lev"]), + "lev": xr.DataArray(np.array(TEST_DATA_LAYOUT.depth_levels), dims=["lev"]), "time": xr.DataArray(time, dims=["time"]), } return coords @@ -322,7 +336,7 @@ def backend(request) -> TrainBackendConfig: return request.param -def _uncached_data_source(name: str) -> DataSource: +def _uncached_data_source(name: str) -> CanonicalSource: match name: case "mock": time_range = xr.cftime_range( @@ -341,22 +355,22 @@ def _uncached_data_source(name: str) -> DataSource: vars_3d = { f"{var}_{lev}": dims.encode(len(vars_2d) + i + j * 10) for i, var in enumerate(["so", "thetao", "uo", "vo"]) - for j, lev in enumerate(TEST_DATASET_SPEC.depth_i_levels) + for j, lev in enumerate(TEST_DATA_LAYOUT.depth_i_levels) } # Mask with a binary circle. masks = { f"mask_{lev}": xr.DataArray( np.where(normal > 0.5**lev, 1, 0), dims=["lat", "lon"] ) - for lev in range(len(TEST_DATASET_SPEC.depth_i_levels)) + for lev in range(len(TEST_DATA_LAYOUT.depth_i_levels)) } ds = xr.Dataset(vars_2d | vars_3d | masks, coords=coords) - return DataSource.from_datasets( + return CanonicalSource.from_datasets( ds, ds.mean(), ds.std(), - dataset_spec=TEST_DATASET_SPEC, + data_layout=TEST_DATA_LAYOUT, name=name, prognostic_var_names=list(vars_3d.keys()), boundary_var_names=list(vars_2d.keys()), @@ -397,7 +411,7 @@ def _fmtl(lev: float) -> str: vars_3d = { f"{var}_lev_{_fmtl(lev)}": dims.encode(len(vars_2d) + i + j * 10) for i, var in enumerate(var_names_3d) - for j, lev in enumerate(TEST_DATASET_SPEC.depth_levels) + for j, lev in enumerate(TEST_DATA_LAYOUT.depth_levels) } # zos is an edge case 3d var. @@ -416,14 +430,15 @@ def _fmtl(lev: float) -> str: data = compact_dataset(data) means = compact_dataset(means) stds = compact_dataset(stds) - prognostic_var_names = var_names_3d boundary_var_names = var_names_2d - return DataSource.from_datasets( + data, means, stds = canonicalize_mock_om4(data, means, stds) + + return CanonicalSource.from_datasets( data=data, means=means, stds=stds, - dataset_spec=TEST_DATASET_SPEC, + data_layout=TEST_DATA_LAYOUT, name=name, prognostic_var_names=prognostic_var_names, boundary_var_names=boundary_var_names, @@ -432,8 +447,10 @@ def _fmtl(lev: float) -> str: raise ValueError(f"Unknown data source: {name}.") -def _maybe_read_cache(cache_root: pathlib.Path, cache_name: str) -> DataSource | None: - """Open a cached DataSource from a cache directory if it exists. +def _maybe_read_cache( + cache_root: pathlib.Path, cache_name: str +) -> CanonicalSource | None: + """Open a cached CanonicalSource from a cache directory if it exists. The caller must ensure concurrent processes/threads do not change this cache. """ @@ -446,33 +463,21 @@ def _maybe_read_cache(cache_root: pathlib.Path, cache_name: str) -> DataSource | boundary_vars = [ str(v) for v in data.data_vars - if v in TEST_FULL_DATASET_SPEC.boundary_var_names + if v in TEST_FULL_DATA_LAYOUT.boundary_var_names ] - if _is_compact(data, means, stds): - prognostic_var_names: list[str] = [] - for var in data.data_vars: - if var in boundary_vars or "mask" in var: - continue - if "lev" in data[var].dims: - prognostic_var_names.extend( - f"{var}_{i}" for i in range(len(data.lev)) - ) - else: - prognostic_var_names.append(str(var)) - else: - prognostic_var_names = _om4_canonical_var_names( - str(v) - for v in data.data_vars - if v not in boundary_vars and "mask" not in v - ) + prognostic_var_names = [ + var + for var in TEST_FULL_DATA_LAYOUT.prognostic_var_names + if var in data.data_vars + ] - return DataSource.from_datasets( + return CanonicalSource.from_datasets( name=cache_name, data=data, means=means, stds=stds, - dataset_spec=TEST_DATASET_SPEC, + data_layout=TEST_DATA_LAYOUT, prognostic_var_names=prognostic_var_names, boundary_var_names=boundary_vars, ) @@ -480,28 +485,30 @@ def _maybe_read_cache(cache_root: pathlib.Path, cache_name: str) -> DataSource | return None -def _write_cache(cache_root: pathlib.Path, data_source: DataSource) -> None: - """Write a DataSource to a new cache directory. +def _write_cache(cache_root: pathlib.Path, data_source: CanonicalSource) -> None: + """Write a CanonicalSource to a new cache directory. The caller must ensure concurrent processes/threads do not read or write to this cache while this function is running. """ cache = cache_root / data_source.name + data, means, stds = data_source._xarray_datasets_for_testing() # Turn off compression! Our training datasets currently have compression turned off. # See https://github.com/m2lines/samudra/blob/main/samudra/__main__.py#L240 assert not (dz := cache / "data.zarr").exists(), "Data already exists in cache" - data_source.data.to_zarr( - dz, encoding={dv: {"compressor": None} for dv in data_source.data.data_vars} + data.to_zarr( + dz, + encoding={dv: {"compressor": None} for dv in data.data_vars}, ) assert not (dm := cache / "means.nc").exists(), "Means already exists in cache" - data_source.means.to_netcdf(dm) + means.to_netcdf(dm) assert not (ds := cache / "stds.nc").exists(), "Stds already exists in cache" - data_source.stds.to_netcdf(ds) + stds.to_netcdf(ds) @pytest.fixture(scope="session", params=["mock", "mock-om4", "compact"]) -def data_source(request, pytestconfig) -> DataSource: +def data_source(request, pytestconfig) -> CanonicalSource: """Returns remote and in-memory `xarray.Dataset`s for tests.""" our_cache_dir = cache_dir(pytestconfig) data_type = request.param @@ -532,7 +539,7 @@ def unique_test_name(config_name: str, pytestconfig: pytest.Config) -> str: @pytest.fixture(scope="function") def train_config( - data_source: DataSource, + data_source: CanonicalSource, pytestconfig: pytest.Config, config_name: str, backend: TrainBackendConfig, @@ -583,23 +590,23 @@ def trainer_pair( @pytest.fixture -def dummy_src(): +def dummy_source(): h, w = 4, 8 - coords = {"lev": [0], "lat": np.arange(h), "lon": np.arange(w)} + coords = {"lat": np.arange(h), "lon": np.arange(w)} data = xr.Dataset( { - "thetao": (("lev", "lat", "lon"), np.zeros((1, h, w))), + "thetao_0": (("lat", "lon"), np.zeros((h, w))), "hfds": (("lat", "lon"), np.zeros((h, w))), }, coords=coords, ) masks = Masks(torch.ones(1, h, w), torch.ones(h, w)) - src = DataSource( + source = CanonicalSource.from_canonical_datasets( name="dummy", data=data, means=data.mean(dim=["lat", "lon"]), stds=data.std(dim=["lat", "lon"]), masks=masks, - dataset_spec=TEST_DATASET_SPEC, + data_layout=TEST_DATA_LAYOUT, ) - yield src + yield source diff --git a/tests/llc_fixtures.py b/tests/llc_fixtures.py index a437e8e0..fceed45f 100644 --- a/tests/llc_fixtures.py +++ b/tests/llc_fixtures.py @@ -7,13 +7,13 @@ import numpy as np import xarray as xr -from samudra.constants import build_llc_spec +from samudra.constants import build_llc_layout def raw_llc_datasets(n_time: int = 3) -> tuple[xr.Dataset, xr.Dataset, xr.Dataset]: """A mock LLC dataset with the same structure but (much) smaller.""" n_face = 2 - n_lev = len(build_llc_spec().depth_i_levels) + n_lev = len(build_llc_layout().depth_i_levels) n_j = 4 n_i = 5 times = np.arange( diff --git a/tests/test_canonical_dataset.py b/tests/test_canonical_dataset.py new file mode 100644 index 00000000..1339b2dc --- /dev/null +++ b/tests/test_canonical_dataset.py @@ -0,0 +1,257 @@ +# SPDX-FileCopyrightText: 2026 Samudra Authors +# +# SPDX-License-Identifier: Apache-2.0 + +from pathlib import Path + +import cftime +import numpy as np +import torch +import xarray as xr + +from samudra.config import JulianDate, Om4TimeConfig +from samudra.datasets import InferenceDataset, TorchTrainDataset +from samudra.utils.data import CanonicalReadRequest, CanonicalSource +from tests.conftest import TEST_DATA_LAYOUT, canonicalize_mock_om4 + + +def _equivalent_om4_sources( + compact_store_root: Path | None = None, +) -> tuple[CanonicalSource, CanonicalSource]: + time = xr.CFTimeIndex( + [cftime.datetime(2000, 1, day, 12, calendar="julian") for day in range(1, 9)] + ) + lat = np.array([-1.0, 1.0]) + lon = np.array([0.5, 1.5, 2.5]) + levels = len(TEST_DATA_LAYOUT.depth_levels) + shape = (len(time), levels, len(lat), len(lon)) + depth_data = np.arange(np.prod(shape), dtype=np.float32).reshape(shape) + surface_data = np.arange(len(time) * len(lat) * len(lon), dtype=np.float32).reshape( + len(time), len(lat), len(lon) + ) + wetmask = np.ones((levels, len(lat), len(lon)), dtype=bool) + wetmask[2, 0, 0] = False + + compact_data = xr.Dataset( + { + "so": (("time", "lev", "lat", "lon"), depth_data), + "zos": (("time", "lat", "lon"), surface_data + 1), + "hfds": (("time", "lat", "lon"), surface_data + 2), + "wetmask": ( + ("lev", "lat", "lon"), + wetmask, + ), + }, + coords={"time": time, "lev": np.arange(levels), "lat": lat, "lon": lon}, + ) + compact_means = compact_data[["so", "zos", "hfds"]].mean(("time", "lat", "lon")) + compact_stds = compact_data[["so", "zos", "hfds"]].std(("time", "lat", "lon")) + + flat_vars: dict[str, tuple[tuple[str, ...], np.ndarray]] = { + **{ + f"so_{level}": ( + ("time", "lat", "lon"), + depth_data[:, level], + ) + for level in range(levels) + }, + **{ + f"mask_{level}": (("lat", "lon"), wetmask[level]) for level in range(levels) + }, + "zos": (("time", "lat", "lon"), surface_data + 1), + "hfds": (("time", "lat", "lon"), surface_data + 2), + } + flat_data = xr.Dataset(flat_vars, coords={"time": time, "lat": lat, "lon": lon}) + data_channels = [f"so_{level}" for level in range(levels)] + ["zos", "hfds"] + flat_means = flat_data[data_channels].mean(("time", "lat", "lon")) + flat_stds = flat_data[data_channels].std(("time", "lat", "lon")) + + prognostic = ["so_0", "so_2", "zos"] + boundary = ["hfds"] + flat_canonical = canonicalize_mock_om4(flat_data, flat_means, flat_stds) + flat = CanonicalSource.from_datasets( + *flat_canonical, + name="flat", + data_layout=TEST_DATA_LAYOUT, + prognostic_var_names=prognostic, + boundary_var_names=boundary, + ) + if compact_store_root is None: + compact_canonical = canonicalize_mock_om4( + compact_data, compact_means, compact_stds + ) + compact = CanonicalSource.from_datasets( + *compact_canonical, + name="compact", + data_layout=TEST_DATA_LAYOUT, + prognostic_var_names=prognostic, + boundary_var_names=boundary, + ) + else: + data_path = compact_store_root / "compact-data.zarr" + means_path = compact_store_root / "compact-means.zarr" + stds_path = compact_store_root / "compact-stds.zarr" + compact_data.to_zarr(data_path, mode="w", consolidated=True) + compact_means.to_zarr(means_path, mode="w", consolidated=True) + compact_stds.to_zarr(stds_path, mode="w", consolidated=True) + compact_canonical = canonicalize_mock_om4( + xr.open_zarr(data_path, chunks={}), + xr.open_zarr(means_path, chunks={}), + xr.open_zarr(stds_path, chunks={}), + ) + compact = CanonicalSource.from_datasets( + *compact_canonical, + data_layout=TEST_DATA_LAYOUT, + prognostic_var_names=prognostic, + boundary_var_names=boundary, + ) + return flat, compact + + +def test_flat_and_compact_om4_have_identical_canonical_contract() -> None: + flat, compact = _equivalent_om4_sources() + + assert flat.channels == compact.channels == ("so_0", "so_2", "zos", "hfds") + np.testing.assert_array_equal(flat.time, compact.time) + torch.testing.assert_close(flat.resolution[0], compact.resolution[0]) + torch.testing.assert_close(flat.resolution[1], compact.resolution[1]) + np.testing.assert_allclose( + flat.statistics(flat.channels).mean, + compact.statistics(compact.channels).mean, + ) + np.testing.assert_allclose( + flat.statistics(flat.channels).std, + compact.statistics(compact.channels).std, + ) + torch.testing.assert_close(flat.masks.prognostic, compact.masks.prognostic) + torch.testing.assert_close(flat.masks.boundary, compact.masks.boundary) + + indices = np.array([[0, 2], [1, 3]], dtype=np.int64) + np.testing.assert_allclose( + flat.read(indices, flat.channels), compact.read(indices, compact.channels) + ) + + +def test_compact_zarr_has_the_same_canonical_contract(tmp_path: Path) -> None: + flat, compact = _equivalent_om4_sources(tmp_path) + indices = np.array([[0, 2], [1, 3]], dtype=np.int64) + + assert compact.channels == flat.channels + np.testing.assert_array_equal(compact.time, flat.time) + np.testing.assert_allclose( + compact.statistics(compact.channels).mean, + flat.statistics(flat.channels).mean, + ) + np.testing.assert_allclose( + compact.statistics(compact.channels).std, + flat.statistics(flat.channels).std, + ) + np.testing.assert_allclose( + compact.read(indices, compact.channels), flat.read(indices, flat.channels) + ) + + +def test_channel_request_and_time_slice_are_immutable() -> None: + source, _ = _equivalent_om4_sources() + original_channels = source.channels + original_time = source.time.copy() + + requested_channels = ("hfds", "so_2") + sliced = source.slice_time( + Om4TimeConfig(start=JulianDate("2000-01-02"), end=JulianDate("2000-01-03")) + ) + + assert source.channels == original_channels + np.testing.assert_array_equal(source.time, original_time) + assert sliced.channels == source.channels + assert sliced.time.size == 2 + np.testing.assert_allclose( + sliced.read(np.array([0]), requested_channels), + source.read(np.array([1]), requested_channels), + ) + + +def test_read_request_owns_immutable_integer_indices() -> None: + indices = np.array([0, 2], dtype=np.int32) + request = CanonicalReadRequest(indices, ("so_0",)) + indices[0] = 1 + + assert request.channels == ("so_0",) + assert request.time_indices.dtype == np.int64 + np.testing.assert_array_equal(request.time_indices, [0, 2]) + assert not request.time_indices.flags.writeable + + +def test_source_reader_can_be_replaced_without_mutating_source() -> None: + source, _ = _equivalent_om4_sources() + replacement = source.reader.slice_time( + Om4TimeConfig(start=JulianDate("2000-01-02"), end=JulianDate("2000-01-03")) + ) + replaced = source.with_reader(replacement) + + assert replaced is not source + assert replaced.reader is replacement + assert source.time.size == 8 + assert replaced.time.size == 2 + assert replaced.channels == source.channels + + +def test_flat_and_compact_cpu_training_and_inference_are_identical() -> None: + flat, compact = _equivalent_om4_sources() + prognostic = ["so_0", "so_2", "zos"] + boundary = ["hfds"] + + def train_dataset( + source: CanonicalSource, *, concurrent: bool = False + ) -> TorchTrainDataset: + return TorchTrainDataset( + input_source=source, + label_source=None, + prognostic_var_names=prognostic, + boundary_var_names=boundary, + hist=1, + steps=1, + normalize_before_mask=True, + masked_fill_value=0.0, + concurrent_compute_=concurrent, + ) + + flat_train = train_dataset(flat) + compact_train = train_dataset(compact) + flat_raw = flat_train[0] + compact_raw = compact_train[0] + for flat_step, compact_step in zip(flat_raw.steps, compact_raw.steps, strict=True): + for flat_tensor, compact_tensor in zip(flat_step, compact_step, strict=True): + torch.testing.assert_close(flat_tensor, compact_tensor) + flat_batch = flat_train.to_model_batch(flat_raw, torch.device("cpu")) + compact_batch = compact_train.to_model_batch(compact_raw, torch.device("cpu")) + for flat_tensor, compact_tensor in zip( + flat_batch[0], compact_batch[0], strict=True + ): + torch.testing.assert_close(flat_tensor, compact_tensor) + + # The concurrent CPU path submits whole canonical read requests, not one task + # per compact level. It must retain the same combined-read semantics. + compact_concurrent = train_dataset(compact, concurrent=True)[0] + for compact_tensor, concurrent_tensor in zip( + compact_raw.steps[0], compact_concurrent.steps[0], strict=True + ): + torch.testing.assert_close(compact_tensor, concurrent_tensor) + + def inference_dataset(source: CanonicalSource) -> InferenceDataset: + return InferenceDataset( + source, + prognostic_var_names=prognostic, + boundary_var_names=boundary, + hist=1, + normalize_before_mask=False, + masked_fill_value=-1.0, + long_rollout=False, + ) + + flat_inference = inference_dataset(flat) + compact_inference = inference_dataset(compact) + for flat_tensor, compact_tensor in zip( + flat_inference[0], compact_inference[0], strict=True + ): + torch.testing.assert_close(flat_tensor, compact_tensor) diff --git a/tests/test_config.py b/tests/test_config.py index edbc94fa..c9736ee9 100644 --- a/tests/test_config.py +++ b/tests/test_config.py @@ -75,16 +75,14 @@ def test_data_config_defaults_to_cpu_loading(): assert isinstance(cfg.sources[0], Om4DataSourceConfig) -def test_om4_dataset_config_builds_selected_spec(): +def test_om4_dataset_config_retains_selected_variable_keys(): cfg = om4_source_config( prognostic_vars_key="thetao_1", boundary_vars_key="hfds", ) - spec = cfg.dataset_spec - - assert spec.prognostic_var_names == ["thetao_0"] - assert spec.boundary_var_names == ["hfds"] + assert cfg.prognostic_vars_key == "thetao_1" + assert cfg.boundary_vars_key == "hfds" def test_data_source_time_configs_use_native_types(): @@ -188,7 +186,7 @@ def test_data_config_accepts_llc_dataset_type(): assert isinstance(source, LlcDataSourceConfig) assert source.face == 2 assert isinstance(source.inference_times[0], LlcTimeConfig) - assert source.dataset_spec.prognostic_var_names == ["Theta_0"] + assert source.prognostic_vars_key == "single_1" def test_data_config_rejects_invalid_llc_crop(): @@ -267,39 +265,58 @@ def test_data_config_builds_llc_source_from_local_files(tmp_path): ) container = cfg.build(LocalLocation(path=tmp_path)) - source = container.primary_source - - assert source.dataset_spec.type == "llc" - assert "Theta_0" in source.data.variables - assert "wetmask_0" in source.data.variables - assert "face" not in source.data.dims - assert source.data["Theta_0"].dims == ("time", "lat", "lon") - assert source.data["wetmask_0"].dims == ("lat", "lon") - assert source.data["Theta_0"].shape == (2, 2, 3) - assert np.issubdtype(source.data.time.dtype, np.datetime64) - assert container.train_sources[0].data.sizes["time"] == 2 - assert container.val_sources[0].data.sizes["time"] == 2 + source = container.train_sources[0] + source_data, _, _ = source._xarray_datasets_for_testing() + + assert source.data_layout.prognostic_var_names == ["Theta_0"] + assert "Theta_0" in source_data.variables + assert "mask_0" in source_data.variables + assert "face" not in source_data.dims + assert source_data["Theta_0"].dims == ("time", "lat", "lon") + assert source_data["mask_0"].dims == ("lat", "lon") + assert source_data["Theta_0"].shape == (2, 2, 3) + assert np.issubdtype(source.time.dtype, np.datetime64) + assert container.train_sources[0].time.size == 2 + assert container.val_sources[0].time.size == 2 assert container.inference_source is not None - assert container.inference_source.data.sizes["time"] == 2 + assert container.inference_source.time.size == 2 - sliced = source.slice( + sliced = source.slice_time( LlcTimeConfig( start=np.datetime64("2011-09-10T12:00:00", "ns"), end=np.datetime64("2011-09-11T12:00:00", "ns"), ) ) - assert sliced.data.sizes["time"] == 2 + assert sliced.time.size == 2 -def test_data_config_rejects_multiple_dataset_specs(tmp_path): - cfg = DataConfig( - sources=[ - om4_source_config(prognostic_vars_key="thetao_1"), - om4_source_config(prognostic_vars_key="thermo_dynamic_all"), - ] +def test_data_config_rejects_multiple_data_layouts(tmp_path): + write_raw_llc_datasets(tmp_path) + source = { + "type": "llc", + "train_time": { + "start": "2011-09-10T12:00:00Z", + "end": "2011-09-11T12:00:00Z", + }, + "val_time": { + "start": "2011-09-11T12:00:00Z", + "end": "2011-09-12T12:00:00Z", + }, + "data_location": "data.zarr", + "data_means_location": "means.nc", + "data_stds_location": "stds.nc", + "prognostic_vars_key": "single_1", + } + cfg = DataConfig.model_validate( + { + "sources": [ + source, + source | {"prognostic_vars_key": "all"}, + ] + } ) - with pytest.raises(AssertionError, match="same dataset spec"): + with pytest.raises(ValueError, match="same data layout"): cfg.build(LocalLocation(path=tmp_path)) diff --git a/tests/test_datasets.py b/tests/test_datasets.py index a64ab299..6e6ae60e 100644 --- a/tests/test_datasets.py +++ b/tests/test_datasets.py @@ -22,21 +22,21 @@ from torch.utils.data import ConcatDataset, DataLoader from samudra.config import DataConfig, TrainConfig -from samudra.constants import LoaderVersion +from samudra.constants import DataLayout, LoaderVersion from samudra.datasets import ( + BatchLoader, InferenceDataset, + ModelBatch, TorchTrainDataset, - TrainData, - TrainDataLoader, ) -from samudra.utils.data import DataSource, Masks, Normalize +from samudra.utils.data import BatchPreprocessor, CanonicalSource, Masks from samudra.utils.location import LocalLocation from samudra.utils.multiton import MultitonScope from samudra.utils.samplers import EquivalenceGroupBatchSampler -from samudra.utils.train import collate_raw_train_data +from samudra.utils.train import collate_host_batches from tests.conftest import ( DEFAULT_CONFIG, - TEST_DATASET_SPEC, + TEST_DATA_LAYOUT, DataSourceDims, TrainPair, cache_dir, @@ -45,9 +45,11 @@ @pytest.fixture -def inference_loader_pair(trainer_pair: TrainPair) -> tuple[TrainConfig, DataLoader]: +def inference_loader_pair( + trainer_pair: TrainPair, +) -> tuple[TrainConfig, DataLoader, DataLayout]: cfg, trainer = trainer_pair - return cfg, trainer.inference_loader + return cfg, trainer.inference_loader, trainer.data_layout def coarsen_data(ds: xr.Dataset) -> xr.Dataset: @@ -55,15 +57,22 @@ def coarsen_data(ds: xr.Dataset) -> xr.Dataset: def coarsen_source( - src: DataSource, + source: CanonicalSource, prognostic: list[str], boundary: list[str], -) -> DataSource: +) -> CanonicalSource: # DEFAULT_CONFIG selects thermo_dynamic_5 plus tau_hfds from the larger mock # OM4 source; coarsen only those active variables to avoid extra Zarr reads. - coarsen_input = src.filter(prognostic + boundary, prefix="coarsen-input") - coarsened_src = coarsen_input.map_data(coarsen_data, suffix="half-size") - return dataclasses.replace(coarsened_src, masks=coarsen_masks(src.masks)) + channels = tuple(prognostic + boundary) + data, means, stds = source._xarray_datasets_for_testing() + return CanonicalSource.from_canonical_datasets( + name=f"{source.name}_half-size", + data=coarsen_data(data[list(channels)]), + means=means[list(channels)], + stds=stds[list(channels)], + masks=coarsen_masks(source.masks), + data_layout=source.data_layout, + ) def coarsen_masks(masks: Masks) -> Masks: @@ -94,7 +103,7 @@ def make_loader( version: LoaderVersion | None = None, multiscale: bool = False, shuffle: bool = True, -) -> Generator[DataLoader | TrainDataLoader, None, None]: +) -> Generator[BatchLoader, None, None]: data_config = ( cfg.data if version is None @@ -103,24 +112,22 @@ def make_loader( container = data_config.build( cfg.experiment.resolved_data_root, ) - dataset_spec = container.dataset_spec - prognostic = dataset_spec.prognostic_var_names - boundary = dataset_spec.boundary_var_names + data_layout = container.data_layout + prognostic = data_layout.prognostic_var_names + boundary = data_layout.boundary_var_names version = container.loader_version - src = container.train_sources[0] - if src.is_compact and version != LoaderVersion.OM4_TORCH: - pytest.skip(f"{version} does not support compact data.") - + source = container.train_sources[0] with MultitonScope(): - srcs = [src] + sources = [source] if multiscale: - srcs.append(coarsen_source(src, prognostic, boundary)) + sources.append(coarsen_source(source, prognostic, boundary)) match version: case LoaderVersion.OM4_TORCH: dataset_list = [ TorchTrainDataset( - src=src, + input_source=source, + label_source=None, prognostic_var_names=prognostic, boundary_var_names=boundary, hist=cfg.data.hist, @@ -129,47 +136,50 @@ def make_loader( masked_fill_value=cfg.data.masked_fill_value, stride=stride, ) - for src in srcs + for source in sources for stride in cfg.data_stride ] - data: ConcatDataset = ConcatDataset(dataset_list) - collate_fn = collate_raw_train_data + host_dataset: ConcatDataset = ConcatDataset(dataset_list) + collate_fn = collate_host_batches - # Group datasets by resolution, allowing different strides to batch together. + # Group datasets by resolution, allowing different strides to batch + # together. batch_sampler = EquivalenceGroupBatchSampler.from_datasets( datasets=dataset_list, - group_key=lambda ds: ds.prognostic_src.grid_size, + group_key=lambda ds: tuple( + source.grid_size for source in ds.sources + ), batch_size=cfg.batch_size, drop_last=drop_last, shuffle=shuffle, seed=cfg.experiment.rand_seed, ) - raw_loader = DataLoader( - data, + host_loader = DataLoader( + host_dataset, batch_sampler=batch_sampler, collate_fn=collate_fn, ) - loader = TrainDataLoader(raw_loader, dataset_list, torch.device("cpu")) + loader = BatchLoader(host_loader, dataset_list, torch.device("cpu")) yield loader case _: raise ValueError(f"Unknown loader version: {version}") -def extract_sample_arrays(td: TrainData) -> tuple[np.ndarray, np.ndarray]: - """Extract underlying X, y pairs from TrainData object. +def extract_sample_arrays(batch: ModelBatch) -> tuple[np.ndarray, np.ndarray]: + """Extract underlying X, y pairs from ModelBatch object. X is the channel-concatenated (prognostic + boundary) tensor for parity with the pre-split-API shape checks these tests do. """ - steps = len(td) + steps = len(batch) x_arrays = [] for s in range(steps): - prog, boundary = td.get_input(s) + prog, boundary = batch.get_input(s) x_arrays.append(torch.cat((prog, boundary), dim=1).numpy(force=True)) - y_arrays = [td.get_label(s).numpy(force=True) for s in range(steps)] + y_arrays = [batch.get_label(s).numpy(force=True) for s in range(steps)] return np.stack(x_arrays, axis=0), np.stack(y_arrays, axis=0) @@ -301,15 +311,14 @@ def test_loader__data_shape( train_config.data.hist = history with make_loader(train_config, version=loader_version) as loader: - dataset_spec = train_config.data.sources[0].dataset_spec + dataset = next(iter(loader._datasets.values())) batch_size = train_config.batch_size num_input_timesteps = history + 1 input_var_dim = ( - len(dataset_spec.prognostic_var_names) - + len(dataset_spec.boundary_var_names) + len(dataset.prognostic_var_names) + len(dataset.boundary_var_names) ) * num_input_timesteps - output_var_dim = len(dataset_spec.prognostic_var_names) * num_input_timesteps + output_var_dim = len(dataset.prognostic_var_names) * num_input_timesteps n_samples = calc_num_samples( train_config, @@ -370,15 +379,14 @@ def test_loader__data_shape__across_source_counts( # Keep grouped-sampler ordering deterministic for resolution coverage. shuffle=False, ) as loader: - dataset_spec = train_config.data.sources[0].dataset_spec + dataset = next(iter(loader._datasets.values())) batch_size = train_config.batch_size num_input_timesteps = history + 1 input_var_dim = ( - len(dataset_spec.prognostic_var_names) - + len(dataset_spec.boundary_var_names) + len(dataset.prognostic_var_names) + len(dataset.boundary_var_names) ) * num_input_timesteps - output_var_dim = len(dataset_spec.prognostic_var_names) * num_input_timesteps + output_var_dim = len(dataset.prognostic_var_names) * num_input_timesteps n_samples = calc_num_samples( train_config, @@ -423,16 +431,14 @@ def test_loader__data_shape__across_source_counts( def test_inference__data_shape(inference_loader_pair): - cfg, loader = inference_loader_pair - - dataset_spec = cfg.data.sources[0].dataset_spec + cfg, loader, data_layout = inference_loader_pair batch_size = 1 # Inference always uses batch size 1 hist = cfg.data.hist + 1 input_var_dim = ( - len(dataset_spec.prognostic_var_names) + len(dataset_spec.boundary_var_names) + len(data_layout.prognostic_var_names) + len(data_layout.boundary_var_names) ) * hist - output_var_dim = len(dataset_spec.prognostic_var_names) * hist + output_var_dim = len(data_layout.prognostic_var_names) * hist samples = list(loader) assert len(samples) == 1, ( @@ -460,7 +466,7 @@ def test__data_is_not_zeros(train_config): def test_inference__data_is_not_zero(inference_loader_pair): - cfg, loader = inference_loader_pair + cfg, loader, _ = inference_loader_pair for sample in loader: dataset, n = sample @@ -504,17 +510,17 @@ def assert_equal_samples(original_samples, new_samples): # to a local directory of cached data. @pytest.mark.parametrize("data_source", ["mock-om4"], indirect=True) def test_compact_loader__equals_flat_loader( - data_source: DataSource, pytestconfig: pytest.Config + data_source: CanonicalSource, pytestconfig: pytest.Config ): cache = cache_dir(pytestconfig) default_config = str(pytestconfig.rootpath / "configs" / DEFAULT_CONFIG) - def make_config(src: DataSource): + def make_config(source: CanonicalSource): return TrainConfig.from_yaml_and_cli( [ default_config, "--experiment.data_root", - str(cache / src.name), + str(cache / source.name), ] ) @@ -593,11 +599,12 @@ def _llc_data_config( def _llc_torch_dataset(config: DataConfig, tmp_path) -> TorchTrainDataset: container = config.build(LocalLocation(path=tmp_path)) - dataset_spec = container.dataset_spec + data_layout = container.data_layout return TorchTrainDataset( - src=container.train_sources[0], - prognostic_var_names=dataset_spec.prognostic_var_names, - boundary_var_names=dataset_spec.boundary_var_names, + input_source=container.train_sources[0], + label_source=None, + prognostic_var_names=data_layout.prognostic_var_names, + boundary_var_names=data_layout.boundary_var_names, hist=config.hist, steps=1, normalize_before_mask=config.normalize_before_mask, @@ -613,12 +620,11 @@ def test_llc_train_dataset_loads_raw_zarr_single_channel(tmp_path): val_end="2011-09-16T12:00:00Z", ) torch_dataset = _llc_torch_dataset(config, tmp_path) - raw_batch = collate_raw_train_data([torch_dataset[0], torch_dataset[1]]) + host_batch = collate_host_batches([torch_dataset[0], torch_dataset[1]]) - train_data = torch_dataset.to_train_data(raw_batch, torch.device("cpu")) - prognostic, boundary, label = train_data[0] - src = torch_dataset.prognostic_src - boundary_src = torch_dataset.boundary_src + model_batch = torch_dataset.to_model_batch(host_batch, torch.device("cpu")) + prognostic, boundary, label = model_batch[0] + source = torch_dataset.sources[0] assert len(torch_dataset) == 3 assert prognostic.shape == (2, 2, 3, 4) @@ -629,17 +635,12 @@ def test_llc_train_dataset_loads_raw_zarr_single_channel(tmp_path): assert label[0, 0, 0, 0] == config.masked_fill_value assert boundary[0, 0, 0, 0] == config.masked_fill_value - assert ( - prognostic[0, 0, 0, 1] == src.data["Theta_0"].isel(time=0, lat=0, lon=1).item() - ) - assert ( - prognostic[0, 1, 0, 1] == src.data["Theta_0"].isel(time=1, lat=0, lon=1).item() - ) - assert label[0, 0, 0, 1] == src.data["Theta_0"].isel(time=2, lat=0, lon=1).item() - assert ( - boundary[0, 0, 0, 1] - == boundary_src.data["oceQnet"].isel(time=0, lat=0, lon=1).item() - ) + source_values = source.read(np.arange(3), torch_dataset.prognostic_var_names) + boundary_values = source.read(np.array([0]), torch_dataset.boundary_var_names) + assert prognostic[0, 0, 0, 1] == source_values[0, 0, 0, 1] + assert prognostic[0, 1, 0, 1] == source_values[1, 0, 0, 1] + assert label[0, 0, 0, 1] == source_values[2, 0, 0, 1] + assert boundary[0, 0, 0, 1] == boundary_values[0, 0, 0, 1] def test_llc_train_dataset_loads_all_raw_variable_families(tmp_path): @@ -652,24 +653,23 @@ def test_llc_train_dataset_loads_all_raw_variable_families(tmp_path): ) config.hist = 0 torch_dataset = _llc_torch_dataset(config, tmp_path) - raw_batch = collate_raw_train_data([torch_dataset[0]]) + host_batch = collate_host_batches([torch_dataset[0]]) - train_data = torch_dataset.to_train_data(raw_batch, torch.device("cpu")) - prognostic, boundary, label = train_data[0] - dataset_spec = config.sources[0].dataset_spec + model_batch = torch_dataset.to_model_batch(host_batch, torch.device("cpu")) + prognostic, boundary, label = model_batch[0] + data_layout = torch_dataset.sources[0].data_layout assert len(torch_dataset) == 3 - assert prognostic.shape == (1, len(dataset_spec.prognostic_var_names), 3, 4) - assert boundary.shape == (1, len(dataset_spec.boundary_var_names), 3, 4) - assert label.shape == (1, len(dataset_spec.prognostic_var_names), 3, 4) + assert prognostic.shape == (1, len(data_layout.prognostic_var_names), 3, 4) + assert boundary.shape == (1, len(data_layout.boundary_var_names), 3, 4) + assert label.shape == (1, len(data_layout.prognostic_var_names), 3, 4) - src = torch_dataset.prognostic_src - assert "Salt_50" in src.data.variables - assert "U_0" in src.data.variables - assert "V_0" in src.data.variables - assert "Eta" in src.data.variables - assert "oceTAUX" in torch_dataset.boundary_src.data.variables - assert "oceTAUY" in torch_dataset.boundary_src.data.variables + assert "Salt_50" in torch_dataset.prognostic_var_names + assert "U_0" in torch_dataset.prognostic_var_names + assert "V_0" in torch_dataset.prognostic_var_names + assert "Eta" in torch_dataset.prognostic_var_names + assert "oceTAUX" in torch_dataset.boundary_var_names + assert "oceTAUY" in torch_dataset.boundary_var_names @pytest.fixture @@ -724,23 +724,24 @@ def tiny_dataset_input(normalize_before_mask: bool, masked_fill_value: float): prognostic=wet, boundary=wet_surface, ) - test = DataSource( + test = CanonicalSource.from_canonical_datasets( "test", data, data_mean, data_std, masks=masks, - dataset_spec=TEST_DATASET_SPEC, + data_layout=TEST_DATA_LAYOUT, ) with MultitonScope(): - _ = Normalize( + _ = BatchPreprocessor( test, prognostic_var_names=["prognostic1", "prognostic2"], boundary_var_names=["boundary1", "boundary2"], ) torch_train_dataset = TorchTrainDataset( - src=test, + input_source=test, + label_source=None, prognostic_var_names=prognostic_var_names, boundary_var_names=boundary_var_names, hist=1, @@ -750,7 +751,7 @@ def tiny_dataset_input(normalize_before_mask: bool, masked_fill_value: float): stride=1, ) inference_dataset = InferenceDataset( - src=test, + source=test, prognostic_var_names=prognostic_var_names, boundary_var_names=boundary_var_names, hist=1, @@ -759,14 +760,14 @@ def tiny_dataset_input(normalize_before_mask: bool, masked_fill_value: float): long_rollout=True, ) - # Create a TrainDataLoader wrapper - raw_loader = DataLoader( + # Create a BatchLoader wrapper + host_loader = DataLoader( torch_train_dataset, batch_size=1, - collate_fn=collate_raw_train_data, + collate_fn=collate_host_batches, ) - train_loader = TrainDataLoader( - raw_loader, [torch_train_dataset], torch.device("cpu") + train_loader = BatchLoader( + host_loader, [torch_train_dataset], torch.device("cpu") ) yield train_loader, inference_dataset @@ -847,7 +848,7 @@ def bench(): "data_source,config_name", [("mock", DEFAULT_CONFIG)], indirect=True ) def test_profile__inference_loader__1gb(inference_loader_pair, benchmark): - cfg, loader = inference_loader_pair + cfg, loader, _ = inference_loader_pair def bench(): for sample in loader: diff --git a/tests/test_gradient_detaching.py b/tests/test_gradient_detaching.py index 19317d13..2bd86d0d 100644 --- a/tests/test_gradient_detaching.py +++ b/tests/test_gradient_detaching.py @@ -8,11 +8,11 @@ import xarray as xr from samudra.config import SamudraConfig, UNetBackboneConfig -from samudra.datasets import TrainData -from samudra.utils.ctx import GridContext -from samudra.utils.data import DataSource, Masks +from samudra.datasets import ModelBatch +from samudra.utils.ctx import BatchGrid +from samudra.utils.data import CanonicalSource, Masks from samudra.utils.multiton import MultitonScope -from tests.conftest import TEST_DATASET_SPEC +from tests.conftest import TEST_DATA_LAYOUT @pytest.fixture(params=[0, 1, 2]) @@ -30,32 +30,31 @@ def _create_model_helper(gradient_detach_interval: int): # Set up minimal data structures needed by Samudra h, w = 8, 8 coords = { - "lev": [0], "lat": (["y"], np.linspace(-90, 90, h)), "lon": (["x"], np.linspace(-180, 180, w)), } data = xr.Dataset( { - "thetao": (["lev", "y", "x"], np.random.randn(1, h, w)), + "thetao_0": (["y", "x"], np.random.randn(h, w)), "hfds": (["y", "x"], np.random.randn(h, w)), }, coords=coords, ) ones = xr.Dataset( { - "thetao": (["lev", "y", "x"], np.ones((1, h, w))), + "thetao_0": (["y", "x"], np.ones((h, w))), "hfds": (["y", "x"], np.ones((h, w))), }, coords=coords, ) masks = Masks(torch.ones(h, w), torch.ones(h, w)) - src = DataSource( + source = CanonicalSource.from_canonical_datasets( name="dummy", data=data, means=data, stds=ones, masks=masks, - dataset_spec=TEST_DATASET_SPEC, + data_layout=TEST_DATA_LAYOUT, ) # Create Samudra model with the specified gradient_detach_interval @@ -72,23 +71,21 @@ def _create_model_helper(gradient_detach_interval: int): boundary_channels=1, out_channels=1, hist=1, - srcs=[src], + grid_sizes=[source.grid_size], ) - # Create TrainData compatible with model dimensions. + # Create ModelBatch compatible with model dimensions. # in_channels=2 splits into 1 prognostic + 1 boundary channel. - train_data = TrainData( - num_prognostic_channels=1, - num_boundary_channels=1, - ctx=GridContext(masks.prognostic, src.resolution, src.resolution), + model_batch = ModelBatch( + BatchGrid(masks.prognostic, source.resolution, source.resolution) ) for step in range(4): prog_tensor = torch.randn(1, 1, h, w, requires_grad=True) boundary_tensor = torch.randn(1, 1, h, w, requires_grad=True) label_tensor = torch.randn(1, 1, h, w) - train_data.append(prog_tensor, boundary_tensor, label_tensor) + model_batch.append(prog_tensor, boundary_tensor, label_tensor) - return model, train_data + return model, model_batch return _create_model_helper @@ -96,10 +93,10 @@ def _create_model_helper(gradient_detach_interval: int): def test_samudra_forward_pass(create_samudra_model, gradient_detach_interval): """Test Samudra forward pass with various gradient detaching intervals.""" interval, interval_desc = gradient_detach_interval - model, train_data = create_samudra_model(interval) + model, model_batch = create_samudra_model(interval) loss_fn = torch.nn.MSELoss() - loss = model(train_data, loss_fn=loss_fn) + loss = model(model_batch, loss_fn=loss_fn) assert not torch.isnan(loss), ( f"Loss is NaN for interval={interval} ({interval_desc})" ) @@ -111,11 +108,11 @@ def test_samudra_forward_pass(create_samudra_model, gradient_detach_interval): def test_samudra_backward_pass(create_samudra_model, gradient_detach_interval): """Test Samudra backward pass with various gradient detaching intervals.""" interval, interval_desc = gradient_detach_interval - model, train_data = create_samudra_model(interval) + model, model_batch = create_samudra_model(interval) loss_fn = torch.nn.MSELoss() # Forward pass - loss = model(train_data, loss_fn=loss_fn) + loss = model(model_batch, loss_fn=loss_fn) # Backward pass loss.backward() diff --git a/tests/test_positional_channels.py b/tests/test_positional_channels.py index 8eab4af8..c718f59c 100644 --- a/tests/test_positional_channels.py +++ b/tests/test_positional_channels.py @@ -5,15 +5,15 @@ import torch from samudra.config import SamudraConfig, UNetBackboneConfig -from samudra.utils.ctx import GridContext -from samudra.utils.data import DataSource +from samudra.utils.ctx import BatchGrid +from samudra.utils.data import CanonicalSource -def test_positional_parameters_update(dummy_src: DataSource): +def test_positional_parameters_update(dummy_source: CanonicalSource): """Verify that positional parameters can learn something in a tiny example.""" - src = dummy_src - h, w = src.grid_size - masks = src.masks + source = dummy_source + h, w = source.grid_size + masks = source.masks # Create the model itself with learned positional embeddings config = SamudraConfig( @@ -29,7 +29,7 @@ def test_positional_parameters_update(dummy_src: DataSource): boundary_channels=1, out_channels=1, hist=0, - srcs=[src], + grid_sizes=[source.grid_size], ) # Verify we have created the positional embeddings @@ -48,7 +48,7 @@ def test_positional_parameters_update(dummy_src: DataSource): out = model.forward_once( prog, boundary, - GridContext(masks.prognostic, src.resolution, src.resolution), + BatchGrid(masks.prognostic, source.resolution, source.resolution), ) loss = out.sum() loss.backward() diff --git a/tests/test_samplers.py b/tests/test_samplers.py index e100cd02..cdc28d94 100644 --- a/tests/test_samplers.py +++ b/tests/test_samplers.py @@ -18,8 +18,6 @@ class MockDataset: def __init__(self, size: int, grid_size: GridSize = (100, 100)): self._size = size - self.input_src = self - self.label_src = self self.grid_size = grid_size def __len__(self): diff --git a/tests/test_samudra_mini.py b/tests/test_samudra_mini.py index d3f7edf4..1379890f 100644 --- a/tests/test_samudra_mini.py +++ b/tests/test_samudra_mini.py @@ -8,7 +8,7 @@ from test_encoder import make_resolution # type: ignore from samudra.models.samudra_mini import SamudraMini -from samudra.utils.ctx import GridContext +from samudra.utils.ctx import BatchGrid def make_perceiver_io( @@ -26,11 +26,11 @@ def make_perceiver_io( ) -def make_ctx(out_channels: int, H: int, W: int) -> GridContext: +def make_ctx(out_channels: int, H: int, W: int) -> BatchGrid: mask = torch.ones(out_channels, H, W, dtype=torch.bool) dummy = torch.randn(1, 1, H, W) res = make_resolution(dummy) - return GridContext(mask, res, res) + return BatchGrid(mask, res, res) def make_model(query_chunk_size: int | None) -> SamudraMini: diff --git a/tests/test_stepper.py b/tests/test_stepper.py index 456d4dda..11a1329b 100644 --- a/tests/test_stepper.py +++ b/tests/test_stepper.py @@ -7,14 +7,13 @@ import torch import xarray as xr -from samudra.constants import TensorMap -from samudra.datasets import InferenceDataset, TrainData +from samudra.datasets import InferenceDataset, ModelBatch from samudra.models.base import BaseModel from samudra.stepper import validate_batch -from samudra.utils.ctx import GridContext -from samudra.utils.data import DataSource, Normalize +from samudra.utils.ctx import BatchGrid +from samudra.utils.data import BatchPreprocessor, CanonicalSource from samudra.utils.multiton import MultitonScope -from tests.conftest import TEST_DATASET_SPEC +from tests.conftest import TEST_DATA_LAYOUT, canonicalize_mock_om4 @pytest.fixture @@ -25,7 +24,7 @@ def inf_data_init(hist: int): lons = 1 total_time_steps = 100 - tensor_map = TensorMap(dataset_spec=TEST_DATASET_SPEC) + data_layout = TEST_DATA_LAYOUT # Even thetao, odd hfds for every time step # Ex, timestep 0: thetao = 0, hfds = 1 @@ -58,32 +57,33 @@ def inf_data_init(hist: int): }, coords={ "time": np.arange(total_time_steps), - "lev": list(TEST_DATASET_SPEC.depth_levels), + "lev": list(TEST_DATA_LAYOUT.depth_levels), "lat": np.arange(lats), "lon": np.arange(lons), }, ) data_mean: xr.Dataset = data.mean() * 0.0 data_std: xr.Dataset = data.std() * 0.0 + 1.0 - val = DataSource.from_datasets( + data, data_mean, data_std = canonicalize_mock_om4(data, data_mean, data_std) + val = CanonicalSource.from_datasets( data, data_mean, data_std, - dataset_spec=TEST_DATASET_SPEC, + data_layout=TEST_DATA_LAYOUT, name="test-data", - prognostic_var_names=tensor_map.prognostic_var_names, - boundary_var_names=tensor_map.boundary_var_names, + prognostic_var_names=data_layout.prognostic_var_names, + boundary_var_names=data_layout.boundary_var_names, ) - _ = Normalize( + _ = BatchPreprocessor( val, - prognostic_var_names=tensor_map.prognostic_var_names, - boundary_var_names=tensor_map.boundary_var_names, + prognostic_var_names=data_layout.prognostic_var_names, + boundary_var_names=data_layout.boundary_var_names, ) inference_dataset = InferenceDataset( val, - tensor_map.prognostic_var_names, - tensor_map.boundary_var_names, + data_layout.prognostic_var_names, + data_layout.boundary_var_names, hist, normalize_before_mask=True, masked_fill_value=0.0, @@ -101,7 +101,7 @@ def forward_once( self, prognostic: torch.Tensor, boundary: torch.Tensor, - ctx: GridContext, + ctx: BatchGrid, ): # Exercises the two streams independently: scale prog, add the last # boundary channel. @@ -112,15 +112,15 @@ class ConstantResidualModel(BaseModel): def __init__(self, *args, **kwargs): super().__init__(*args, **kwargs) - def forward_once(self, prognostic, boundary, ctx: GridContext): + def forward_once(self, prognostic, boundary, ctx: BatchGrid): return torch.ones_like(prognostic) def test_validate_batch_uses_absolute_predictions_for_residual_models(): wet = torch.ones((1, 1, 1, 1), dtype=torch.bool) grid = torch.zeros(1) - ctx = GridContext(wet, (grid, grid), (grid, grid)) - batch = TrainData(num_prognostic_channels=1, num_boundary_channels=1, ctx=ctx) + ctx = BatchGrid(wet, (grid, grid), (grid, grid)) + batch = ModelBatch(ctx) prog_input = torch.tensor([[[[10.0]]]]) boundary_input = torch.tensor([[[[5.0]]]]) label = torch.tensor([[[[11.0]]]]) @@ -216,11 +216,11 @@ def test_inference_rollout(inf_data_init, hist, num_steps): model.eval() initial_prognostic = inference_dataset.initial_prognostic - IO = model.inference( + inference_output = model.inference( inference_dataset, initial_prognostic, num_steps=num_steps, epoch=0 ) - prediction = IO.prediction - target = IO.target + prediction = inference_output.prediction + target = inference_output.target assert prediction.shape == target.shape @@ -283,7 +283,11 @@ def test_inference_rollout_methods(inf_data_init, hist, merge_step): pred = model.forward_once( prog_tensor, boundary_tensor, - GridContext(wet, inference_dataset.input_res, inference_dataset.input_res), + BatchGrid( + wet, + inference_dataset.input_resolution, + inference_dataset.input_resolution, + ), ) assert pred.shape == (1, num_prognostic_channels, 1, 1) expected_pred = torch.tensor( diff --git a/tests/test_train_progress.py b/tests/test_train_progress.py index 86e0a390..7d5da093 100644 --- a/tests/test_train_progress.py +++ b/tests/test_train_progress.py @@ -4,12 +4,12 @@ import torch -from samudra.datasets import TrainData -from samudra.utils.ctx import GridContext +from samudra.datasets import ModelBatch +from samudra.utils.ctx import BatchGrid from samudra.utils.train_progress import TrainBatchProgress, TrainProgress -def make_train_data( +def make_model_batch( *, batch_size: int = 2, input_channels: int = 3, @@ -18,8 +18,8 @@ def make_train_data( input_grid: tuple[int, int] = (3, 4), output_grid: tuple[int, int] = (5, 6), num_model_steps: int = 2, -) -> TrainData: - ctx = GridContext( +) -> ModelBatch: + ctx = BatchGrid( label_mask=torch.ones(output_channels, *output_grid, dtype=torch.bool), input_resolution_cpu=(torch.arange(input_grid[0]), torch.arange(input_grid[1])), output_resolution_cpu=( @@ -27,14 +27,14 @@ def make_train_data( torch.arange(output_grid[1]), ), ) - train_data = TrainData(input_channels, boundary_channels, ctx) + batch = ModelBatch(ctx) for _ in range(num_model_steps): - train_data.append( + batch.append( torch.zeros(batch_size, input_channels, *input_grid), torch.zeros(batch_size, boundary_channels, *input_grid), torch.zeros(batch_size, output_channels, *output_grid), ) - return train_data + return batch def test_train_batch_progress_counts_global_training_units(): @@ -44,7 +44,7 @@ def test_train_batch_progress_counts_global_training_units(): input_grid = (3, 4) output_grid = (5, 6) num_model_steps = 2 - train_data = make_train_data( + batch = make_model_batch( batch_size=batch_size, output_channels=output_channels, input_grid=input_grid, @@ -52,7 +52,7 @@ def test_train_batch_progress_counts_global_training_units(): num_model_steps=num_model_steps, ) - progress = TrainBatchProgress.from_train_data(train_data, world_size) + progress = TrainBatchProgress.from_model_batch(batch, world_size) assert progress.sample_windows == batch_size * world_size assert progress.model_examples == batch_size * world_size * num_model_steps @@ -76,8 +76,8 @@ def test_train_batch_progress_counts_global_training_units(): def test_train_batch_progress_throughput_metrics_use_batch_seconds(): - train_data = make_train_data(batch_size=1, output_channels=1, output_grid=(2, 3)) - progress = TrainBatchProgress.from_train_data(train_data, world_size=2) + batch = make_model_batch(batch_size=1, output_channels=1, output_grid=(2, 3)) + progress = TrainBatchProgress.from_model_batch(batch, world_size=2) progress.batch_seconds = 0.5 metrics = progress.to_throughput_metrics() @@ -91,7 +91,7 @@ def test_train_batch_progress_throughput_metrics_use_batch_seconds(): def test_train_progress_batch_context_records_elapsed_progress(monkeypatch): - train_data = make_train_data(batch_size=1, output_channels=1, output_grid=(2, 3)) + batch = make_model_batch(batch_size=1, output_channels=1, output_grid=(2, 3)) progress = TrainProgress() times = iter([10.0, 12.5]) monkeypatch.setattr( @@ -99,7 +99,7 @@ def test_train_progress_batch_context_records_elapsed_progress(monkeypatch): ) with progress.batch( - train_data, world_size=2, device=torch.device("cpu") + batch, world_size=2, device=torch.device("cpu") ) as batch_progress: batch_progress.optimizer_stepped = True diff --git a/tests/test_trainer.py b/tests/test_trainer.py index bf4337a7..adf92088 100644 --- a/tests/test_trainer.py +++ b/tests/test_trainer.py @@ -13,7 +13,7 @@ from samudra.config import CpuDataLoadingConfig, DynamicLossConfig, TrainConfig from samudra.models.base import BaseModel from samudra.train import Trainer, should_log_validation_images -from samudra.utils.ctx import GridContext +from samudra.utils.ctx import BatchGrid from samudra.utils.loss import DynamicLoss from samudra.utils.multiton import MultitonScope from tests.conftest import DEFAULT_CONFIG, SAMUDRA_MULTI_CONFIG, TrainPair @@ -247,10 +247,10 @@ def test_checkpoint_inference(trainer_pair: TrainPair, caplog): _, trainer = trainer_pair hist = trainer.hist - assert trainer.inference_src is not None - resolution = trainer.inference_src.resolution - wet = trainer.inference_src.masks.prognostic_with_hist(hist) - ctx = GridContext(wet, resolution, resolution).to(trainer.device) + assert trainer.inference_source is not None + resolution = trainer.inference_source.resolution + wet = trainer.inference_source.masks.prognostic_with_hist(hist) + ctx = BatchGrid(wet, resolution, resolution).to(trainer.device) data = trainer.inference_loader.dataset[0] inference_dataset, _num_steps = data prog, boundary, _label = inference_dataset[0] @@ -323,7 +323,7 @@ def test_multiscale_training_validates_primary_source_and_logs_reduced_metrics( assert len(trainer.train_loader._datasets) == 2 assert len(trainer.val_loader._datasets) == 1 val_dataset = next(iter(trainer.val_loader._datasets.values())) - assert val_dataset.prognostic_src.grid_size == trainer.primary_src.grid_size + assert val_dataset.sources[0].grid_size == trainer.primary_source.grid_size class PerfectModel(BaseModel): def __init__(self): @@ -352,9 +352,9 @@ def test_data_loaders_enable_persistent_workers_on_positive_num_workers( assert trainer.mp_context is not None assert trainer.mp_context.get_start_method() == "spawn" - assert trainer.train_loader._dataloader.persistent_workers is True - assert trainer.val_loader._dataloader.persistent_workers is True - assert trainer.inference_src is not None + assert trainer.train_loader._host_loader.persistent_workers is True + assert trainer.val_loader._host_loader.persistent_workers is True + assert trainer.inference_source is not None @pytest.mark.parametrize("backend", ["cpu"], indirect=True) @@ -375,5 +375,5 @@ def test_data_loaders_disable_persistent_workers_when_num_workers_is_zero( trainer.init_data_loaders(cur_step=train_config.steps[0]) assert trainer.mp_context is None - assert trainer.train_loader._dataloader.persistent_workers is False - assert trainer.val_loader._dataloader.persistent_workers is False + assert trainer.train_loader._host_loader.persistent_workers is False + assert trainer.val_loader._host_loader.persistent_workers is False diff --git a/tests/test_utils_data.py b/tests/test_utils_data.py index 5e1d1847..1545842a 100644 --- a/tests/test_utils_data.py +++ b/tests/test_utils_data.py @@ -10,12 +10,11 @@ import xarray as xr from scipy.stats import pearsonr -from samudra.constants import TensorMap, build_llc_spec +from samudra.constants import build_llc_layout from samudra.utils.data import ( - DataSource, + BatchPreprocessor, + CanonicalSource, Masks, - Normalize, - OceanData, compute_anomalies, flatten_masks, get_aggregator_dicts, @@ -31,17 +30,26 @@ _var_without_level, canonicalize_llc_datasets, ) -from tests.conftest import TEST_DATASET_SPEC, TEST_FULL_DATASET_SPEC +from tests.conftest import ( + TEST_DATA_LAYOUT, + TEST_FULL_DATA_LAYOUT, + canonicalize_mock_om4, +) from tests.llc_fixtures import raw_llc_datasets def test_mask_roundtrip(data_source): - data = data_source.data + data, _, _ = data_source._xarray_datasets_for_testing() - unflattened = unflatten_masks(data.copy(), dataset_spec=TEST_DATASET_SPEC) - flattened = flatten_masks(unflattened.copy(), dataset_spec=TEST_DATASET_SPEC) + num_levels = len(TEST_DATA_LAYOUT.depth_levels) + unflattened = unflatten_masks(data.copy(), num_levels=num_levels) + flattened = flatten_masks(unflattened.copy()) + mask_vars = [f"mask_{level}" for level in range(num_levels)] - assert flattened == data, "Assume a safe roundtrip" + xr.testing.assert_equal( + flattened[mask_vars], + data[mask_vars], + ) @pytest.mark.parametrize("data_source", ["mock-om4"], indirect=True) @@ -51,29 +59,34 @@ def test_level_index_vars_roundtrip(data_source): Exercised on the mock OM4 dataset (in ``_`` form) in both orders, so each function is run against the other's real output. """ - spec = TEST_FULL_DATASET_SPEC - ds_idx = data_source.data # OM4 data named _ + data_layout = TEST_FULL_DATA_LAYOUT + ds_idx, _, _ = data_source._xarray_datasets_for_testing() - ds_lev = with_depth_value_vars(ds_idx, spec) + ds_lev = with_depth_value_vars(ds_idx, data_layout) # The inverse actually renamed the 3D vars to the depth-value form. assert any("_lev_" in str(v) for v in ds_lev.variables) # inverse -> forward recovers the index form ... - xr.testing.assert_identical(with_level_index_vars(ds_lev, spec), ds_idx) + xr.testing.assert_identical( + with_level_index_vars(ds_lev, data_layout.depth_levels), ds_idx + ) # ... and forward -> inverse recovers the depth-value form. xr.testing.assert_identical( - with_depth_value_vars(with_level_index_vars(ds_lev, spec), spec), ds_lev + with_depth_value_vars( + with_level_index_vars(ds_lev, data_layout.depth_levels), data_layout + ), + ds_lev, ) @pytest.mark.parametrize("data_source", ["mock-om4"], indirect=True) def test_stack_levels(data_source): """`stack_levels` reassembles flattened OM4 data into depth-stacked form.""" - spec = TEST_FULL_DATASET_SPEC - ds = data_source.data - n = len(spec.depth_levels) + data_layout = TEST_FULL_DATA_LAYOUT + ds, _, _ = data_source._xarray_datasets_for_testing() + n = len(data_layout.depth_levels) - stacked = stack_levels(ds, spec) + stacked = stack_levels(ds, data_layout) # 3D vars gain a `lev` dimension; per-level channels are gone. for base in ["thetao", "so", "uo", "vo"]: @@ -141,7 +154,7 @@ def test_rename_vars(): ) # Apply rename_vars - renamed_ds = with_level_index_vars(ds, dataset_spec=TEST_DATASET_SPEC) + renamed_ds = with_level_index_vars(ds, depth_levels=TEST_DATA_LAYOUT.depth_levels) # Test that variables are renamed correctly assert "so_11" in renamed_ds.variables # 1040.0 is OM4 depth index 11 @@ -178,7 +191,7 @@ def test_rename_vars_invalid_depth(): # Should raise ValueError because 9999.0 is not an OM4 depth level with pytest.raises(ValueError): - with_level_index_vars(ds, dataset_spec=TEST_DATASET_SPEC) + with_level_index_vars(ds, depth_levels=TEST_DATA_LAYOUT.depth_levels) def test_compute_anomalies(): @@ -222,7 +235,7 @@ def test_compute_anomalies(): @pytest.fixture -def normalize_input(): +def preprocessor_input(): # Create test data with mean and std data_mean = xr.Dataset( { @@ -247,70 +260,72 @@ def normalize_input(): # Warning: the 'data' field is not used because this test tries to test # normalization which only needs mean and std. Thus, we set it to `data_mean`. - test = DataSource( + test = CanonicalSource.from_canonical_datasets( "test", data_mean, data_mean, data_std, masks=masks, - dataset_spec=TEST_DATASET_SPEC, + data_layout=TEST_DATA_LAYOUT, ) - normalize = Normalize( + preprocessor = BatchPreprocessor( test, prognostic_var_names=["var_0", "var_1"], boundary_var_names=["var_2"], ) - yield normalize, wet_mask + yield preprocessor, wet_mask -def test_normalize_unnormalize_tensor_prognostic(normalize_input): - normalize, wet_mask = normalize_input - data = torch.randn([1, normalize._prognostic_std_np.shape[0], *wet_mask.shape]) +def test_normalize_unnormalize_tensor_prognostic(preprocessor_input): + preprocessor, wet_mask = preprocessor_input + data = torch.randn([1, preprocessor._prognostic_std_np.shape[0], *wet_mask.shape]) input_data = data * wet_mask - normalized = normalize.normalize_tensor_prognostic(input_data) - unnormalized = normalize.unnormalize_tensor_prognostic(normalized, fill_value=0.0) + normalized = preprocessor.normalize_tensor_prognostic(input_data) + unnormalized = preprocessor.unnormalize_tensor_prognostic( + normalized, fill_value=0.0 + ) assert torch.allclose(input_data, unnormalized) @pytest.mark.parametrize("fill_value", [float("nan"), 0.0]) -def test_unnormalize_prognostic_tensor(normalize_input, fill_value): - normalize, wet_mask = normalize_input - data = torch.randn([1, normalize._prognostic_std_np.shape[0], *wet_mask.shape]) +def test_unnormalize_prognostic_tensor(preprocessor_input, fill_value): + preprocessor, wet_mask = preprocessor_input + data = torch.randn([1, preprocessor._prognostic_std_np.shape[0], *wet_mask.shape]) input_data = data * wet_mask - normalized = normalize.normalize_tensor_prognostic(input_data) - unnormalized = normalize.unnormalize_tensor_prognostic(normalized, fill_value) + normalized = preprocessor.normalize_tensor_prognostic(input_data) + unnormalized = preprocessor.unnormalize_tensor_prognostic(normalized, fill_value) assert (torch.sum(torch.isnan(unnormalized)) > 0) == (math.isnan(fill_value)) @pytest.mark.parametrize("data_source", ["compact"], indirect=True) def test_normalize_compact_mixed_depth_and_surface_stats(data_source): - src = DataSource.from_datasets( - data_source.data, - data_source.means, - data_source.stds, - dataset_spec=TEST_FULL_DATASET_SPEC, + source = CanonicalSource.from_datasets( + *data_source._xarray_datasets_for_testing(), + data_layout=TEST_FULL_DATA_LAYOUT, name="compact-full", - prognostic_var_names=TEST_FULL_DATASET_SPEC.prognostic_var_names, - boundary_var_names=TEST_FULL_DATASET_SPEC.boundary_var_names, + prognostic_var_names=TEST_FULL_DATA_LAYOUT.prognostic_var_names, + boundary_var_names=TEST_FULL_DATA_LAYOUT.boundary_var_names, ) - normalize = Normalize( - src, - prognostic_var_names=TEST_FULL_DATASET_SPEC.prognostic_var_names, - boundary_var_names=TEST_FULL_DATASET_SPEC.boundary_var_names, + preprocessor = BatchPreprocessor( + source, + prognostic_var_names=TEST_FULL_DATA_LAYOUT.prognostic_var_names, + boundary_var_names=TEST_FULL_DATA_LAYOUT.boundary_var_names, ) - num_depth = len(TEST_FULL_DATASET_SPEC.depth_levels) + num_depth = len(TEST_FULL_DATA_LAYOUT.depth_levels) expected_prognostic_channels = 4 * num_depth + 1 assert expected_prognostic_channels == len( - TEST_FULL_DATASET_SPEC.prognostic_var_names + TEST_FULL_DATA_LAYOUT.prognostic_var_names ) - assert normalize._prognostic_mean_np.shape == (expected_prognostic_channels,) - assert normalize._prognostic_std_np.shape == (expected_prognostic_channels,) + assert preprocessor._prognostic_mean_np.shape == (expected_prognostic_channels,) + assert preprocessor._prognostic_std_np.shape == (expected_prognostic_channels,) - lat, lon = src.grid_size + lat, lon = source.grid_size prognostic = torch.zeros(1, expected_prognostic_channels, lat, lon) - assert normalize.normalize_tensor_prognostic(prognostic).shape == prognostic.shape + assert ( + preprocessor.normalize_tensor_prognostic(prognostic).shape == prognostic.shape + ) def test_rename_llc_level_index_vars(): @@ -331,20 +346,20 @@ def test_rename_llc_level_index_vars(): def test_flatten_llc_level_vars(): - raw_data, _, _ = raw_llc_datasets() - data = raw_data[["Theta"]].isel(face=0, drop=True).rename({"k": "lev"}) - dataset_spec = build_llc_spec() + steps, _, _ = raw_llc_datasets() + data = steps[["Theta"]].isel(face=0, drop=True).rename({"k": "lev"}) + data_layout = build_llc_layout() - flattened = _flatten_llc_level_vars(data, dataset_spec=dataset_spec) + flattened = _flatten_llc_level_vars(data, num_levels=len(data_layout.depth_levels)) assert "Theta" not in flattened.data_vars assert set(flattened.data_vars) == { - f"Theta_{level}" for level in dataset_spec.depth_i_levels + f"Theta_{level}" for level in data_layout.depth_i_levels } xr.testing.assert_identical( flattened["Theta_0"], data["Theta"].isel(lev=0, drop=True).rename("Theta_0") ) - last_level = dataset_spec.depth_i_levels[-1] + last_level = data_layout.depth_i_levels[-1] xr.testing.assert_identical( flattened[f"Theta_{last_level}"], data["Theta"].isel(lev=-1, drop=True).rename(f"Theta_{last_level}"), @@ -365,10 +380,10 @@ def test_var_without_level(var_name, expected): def test_canonicalize_llc_datasets_standardizes_layout(): data, means, stds = raw_llc_datasets() - dataset_spec = build_llc_spec(prognostic_vars_key="all", boundary_vars_key="all") + data_layout = build_llc_layout(prognostic_vars_key="all", boundary_vars_key="all") expected_theta_0 = data["Theta"].isel(time=0, face=1, k=0, j=1, i=1).item() - llc_data, llc_means, llc_stds = canonicalize_llc_datasets( + llc_data, llc_means, llc_stds, returned_layout = canonicalize_llc_datasets( data, means, stds, @@ -377,9 +392,12 @@ def test_canonicalize_llc_datasets_standardizes_layout(): i_end=4, j_start=1, j_end=3, - dataset_spec=dataset_spec, + prognostic_vars_key="all", + boundary_vars_key="all", ) + assert returned_layout == data_layout + assert "face" not in llc_data.dims assert "Theta" not in llc_data.variables assert "wetmask" not in llc_data.variables @@ -387,17 +405,17 @@ def test_canonicalize_llc_datasets_standardizes_layout(): assert "Theta_50" in llc_data.variables assert "U_0" in llc_data.variables assert "V_0" in llc_data.variables - assert "wetmask_0" in llc_data.variables + assert "mask_0" in llc_data.variables assert "mask_w_0" in llc_data.variables assert "mask_s_0" in llc_data.variables - assert llc_data["Theta_0"].dims == ("time", "y", "x") - assert llc_data["U_0"].dims == ("time", "y", "x") - assert llc_data["V_0"].dims == ("time", "y", "x") - assert llc_data["wetmask_0"].dims == ("y", "x") - assert llc_data["mask_w_0"].dims == ("y", "x") - assert llc_data["mask_s_0"].dims == ("y", "x") + assert llc_data["Theta_0"].dims == ("time", "lat", "lon") + assert llc_data["U_0"].dims == ("time", "lat", "lon") + assert llc_data["V_0"].dims == ("time", "lat", "lon") + assert llc_data["mask_0"].dims == ("lat", "lon") + assert llc_data["mask_w_0"].dims == ("lat", "lon") + assert llc_data["mask_s_0"].dims == ("lat", "lon") assert llc_data["Theta_0"].shape == (3, 2, 3) - assert llc_data["Theta_0"].isel(time=0, y=0, x=0).item() == expected_theta_0 + assert llc_data["Theta_0"].isel(time=0, lat=0, lon=0).item() == expected_theta_0 assert np.issubdtype(llc_data.time.dtype, np.datetime64) assert "Theta_0" in llc_means.variables assert "Theta_0" in llc_stds.variables @@ -408,7 +426,7 @@ def test_canonicalize_llc_datasets_standardizes_layout(): def test_canonicalize_llc_datasets_selects_requested_vars_from_full_root(): data, means, stds = raw_llc_datasets() - llc_data, _, _ = canonicalize_llc_datasets( + llc_data, _, _, llc_layout = canonicalize_llc_datasets( data, means, stds, @@ -417,14 +435,14 @@ def test_canonicalize_llc_datasets_selects_requested_vars_from_full_root(): i_end=4, j_start=1, j_end=3, - dataset_spec=build_llc_spec(), + prognostic_vars_key="single_1", + boundary_vars_key="single_1", ) - llc_spec = build_llc_spec() expected_vars = { - *(f"Theta_{i}" for i in llc_spec.depth_i_levels), + *(f"Theta_{i}" for i in llc_layout.depth_i_levels), "oceQnet", - *llc_spec.mask_vars, + *(f"mask_{i}" for i in llc_layout.depth_i_levels), } assert expected_vars.issubset(llc_data.data_vars) assert "XG" not in llc_data.data_vars @@ -434,8 +452,8 @@ def test_canonicalize_llc_datasets_selects_requested_vars_from_full_root(): def test_llc_all_variable_masks_use_staggered_masks(): data, means, stds = raw_llc_datasets() - dataset_spec = build_llc_spec(prognostic_vars_key="all", boundary_vars_key="all") - llc_data, llc_means, llc_stds = canonicalize_llc_datasets( + data_layout = build_llc_layout(prognostic_vars_key="all", boundary_vars_key="all") + llc_data, llc_means, llc_stds, returned_layout = canonicalize_llc_datasets( data, means, stds, @@ -444,29 +462,31 @@ def test_llc_all_variable_masks_use_staggered_masks(): i_end=4, j_start=1, j_end=3, - dataset_spec=dataset_spec, + prognostic_vars_key="all", + boundary_vars_key="all", ) + assert returned_layout == data_layout - source = DataSource.from_datasets( + source = CanonicalSource.from_datasets( llc_data, llc_means, llc_stds, - dataset_spec=dataset_spec, - prognostic_var_names=dataset_spec.prognostic_var_names, - boundary_var_names=dataset_spec.boundary_var_names, + data_layout=data_layout, + prognostic_var_names=data_layout.prognostic_var_names, + boundary_var_names=data_layout.boundary_var_names, ) - theta_index = dataset_spec.prognostic_var_names.index("Theta_0") - u_index = dataset_spec.prognostic_var_names.index("U_0") - v_index = dataset_spec.prognostic_var_names.index("V_0") + theta_index = data_layout.prognostic_var_names.index("Theta_0") + u_index = data_layout.prognostic_var_names.index("U_0") + v_index = data_layout.prognostic_var_names.index("V_0") assert bool(source.masks.prognostic[theta_index, 0, 0]) assert not bool(source.masks.prognostic[u_index, 0, 0]) assert not bool(source.masks.prognostic[v_index, 0, 1]) - tau_x_index = dataset_spec.boundary_var_names.index("oceTAUX") - tau_y_index = dataset_spec.boundary_var_names.index("oceTAUY") - qnet_index = dataset_spec.boundary_var_names.index("oceQnet") - assert source.masks.boundary.shape == (len(dataset_spec.boundary_var_names), 2, 3) + tau_x_index = data_layout.boundary_var_names.index("oceTAUX") + tau_y_index = data_layout.boundary_var_names.index("oceTAUY") + qnet_index = data_layout.boundary_var_names.index("oceQnet") + assert source.masks.boundary.shape == (len(data_layout.boundary_var_names), 2, 3) assert not bool(source.masks.boundary[tau_x_index, 0, 0]) assert not bool(source.masks.boundary[tau_y_index, 0, 1]) assert bool(source.masks.boundary[qnet_index, 0, 0]) @@ -479,7 +499,7 @@ def data_init(hist: int): lons = 3 total_time_steps = 100 - tensor_map = TensorMap(dataset_spec=TEST_DATASET_SPEC) + data_layout = TEST_DATA_LAYOUT wet_mask_ = np.array([[1, 0, 1], [0, 1, 0], [1, 0, 1]]) wet_full = np.tile(wet_mask_, (total_time_steps, levels, 1, 1)) @@ -515,39 +535,40 @@ def data_init(hist: int): }, coords={ "time": np.arange(total_time_steps), - "lev": list(TEST_DATASET_SPEC.depth_levels), + "lev": list(TEST_DATA_LAYOUT.depth_levels), "lat": np.arange(lats), "lon": np.arange(lons), }, ) data_mean = data.mean() * 0.0 data_std = data.std() * 0.0 + 1.0 - val = DataSource.from_datasets( + data, data_mean, data_std = canonicalize_mock_om4(data, data_mean, data_std) + val = CanonicalSource.from_datasets( data, data_mean, data_std, - dataset_spec=TEST_DATASET_SPEC, + data_layout=TEST_DATA_LAYOUT, name="test", - prognostic_var_names=tensor_map.prognostic_var_names, - boundary_var_names=tensor_map.boundary_var_names, + prognostic_var_names=data_layout.prognostic_var_names, + boundary_var_names=data_layout.boundary_var_names, ) - normalize = Normalize( + preprocessor = BatchPreprocessor( val, - prognostic_var_names=tensor_map.prognostic_var_names, - boundary_var_names=tensor_map.boundary_var_names, + prognostic_var_names=data_layout.prognostic_var_names, + boundary_var_names=data_layout.boundary_var_names, ) - yield normalize, val.masks.prognostic, tensor_map + yield preprocessor, val.masks.prognostic, data_layout @pytest.mark.parametrize("input_type", ["input", "target"]) @pytest.mark.parametrize("long_rollout", [True, False]) @pytest.mark.parametrize("hist", [0, 1, 2]) def test_get_norm_unnorm_dicts(data_init, input_type, long_rollout, hist): - normalize, wet, tensor_map = data_init + preprocessor, wet, data_layout = data_init - num_prognostic_channels = normalize._prognostic_std_np.shape[0] - num_boundary_channels = normalize._boundary_std_np.shape[0] + num_prognostic_channels = preprocessor._prognostic_std_np.shape[0] + num_boundary_channels = preprocessor._boundary_std_np.shape[0] if input_type == "target": data = torch.randn([1, num_prognostic_channels * (hist + 1), *wet.shape[1:]]) elif input_type == "input": @@ -560,8 +581,8 @@ def test_get_norm_unnorm_dicts(data_init, input_type, long_rollout, hist): ) data_dict, data_unnorm_dict = get_aggregator_dicts( data, - normalize=normalize, - tensor_map=tensor_map, + preprocessor=preprocessor, + data_layout=data_layout, wet=wet, long_rollout=long_rollout, input_type=input_type, @@ -569,76 +590,9 @@ def test_get_norm_unnorm_dicts(data_init, input_type, long_rollout, hist): hist=hist, ) - var_name = tensor_map.prognostic_var_names[0] + var_name = data_layout.prognostic_var_names[0] assert data_dict[var_name].shape == data_unnorm_dict[var_name].shape assert torch.isnan(data_dict[var_name][:, :, 0, 1]).all() assert torch.isnan(data_dict[var_name][:, :, 1, 0]).all() assert torch.isnan(data_dict[var_name][:, :, 1, 2]).all() - - -def test_ocean_data_with_time(): - """Test slicing OceanData across the time dimension.""" - batch, time, var, lat, lon = 2, 5, 3, 4, 6 - data = torch.randn(batch, time, var, lat, lon) - means = torch.tensor([1.0, 2.0, 3.0]) - stds = torch.tensor([0.5, 1.0, 2.0]) - mask = torch.ones(var, dtype=torch.bool) - ocean_data = OceanData(data=data, means=means, stds=stds, mask=mask) - - sliced = ocean_data.with_time(slice(0, 3)) - - assert sliced.data.shape[1] == 3 - assert torch.equal(sliced.data, ocean_data.data[:, 0:3, :, :, :]) - # Other fields should be unchanged - assert torch.equal(sliced.means, ocean_data.means) - assert torch.equal(sliced.stds, ocean_data.stds) - - -@pytest.mark.parametrize("masked_fill_value", [0.0, -1.0]) -def test_ocean_data_normalize_and_mask(masked_fill_value): - """Test that masked positions receive the fill value when normalizing first.""" - batch, time, var, lat, lon = 2, 5, 3, 4, 6 - data = torch.randn(batch, time, var, lat, lon) - data[:, :, :, 0, 0] = float("nan") # Simulate land - means = torch.tensor([1.0, 2.0, 3.0]) - stds = torch.tensor([0.5, 1.0, 2.0]) - mask = torch.ones(var, lat, lon, dtype=torch.bool) - mask[:, 0, 0] = False # Mark as land - - ocean_data = OceanData(data=data, means=means, stds=stds, mask=mask) - result = ocean_data.normalize_and_mask( - normalize_before_mask=True, masked_fill_value=masked_fill_value - ) - - assert result.shape == ocean_data.data.shape - # Masked positions should have the fill value - mask_expanded = mask.unsqueeze(0).unsqueeze(0) - masked_positions = ~mask_expanded.expand_as(result) - assert torch.all(result[masked_positions] == masked_fill_value) - # Valid positions should not be NaN - valid_positions = mask_expanded.expand_as(result) - assert not torch.any(torch.isnan(result[valid_positions])) - - -def test_ocean_data_normalize_and_mask_values(): - """Test that normalization produces expected values with known inputs.""" - batch, time, num_var, lat, lon = 1, 1, 2, 2, 2 - data = torch.tensor( - [[[[[10.0, 10.0], [10.0, 10.0]], [[20.0, 20.0], [20.0, 20.0]]]]] - ) - means = torch.tensor([5.0, 10.0]) - stds = torch.tensor([5.0, 5.0]) - mask = torch.ones(num_var, lat, lon, dtype=torch.bool) - - ocean_data = OceanData(data=data, means=means, stds=stds, mask=mask) - result = ocean_data.normalize_and_mask( - normalize_before_mask=True, masked_fill_value=0.0 - ) - - # Expected: (10 - 5) / 5 = 1.0 for var 0, (20 - 10) / 5 = 2.0 for var 1 - expected_var0 = torch.ones(batch, time, lat, lon) - expected_var1 = torch.ones(batch, time, lat, lon) * 2.0 - - assert torch.allclose(result[:, :, 0, :, :], expected_var0) - assert torch.allclose(result[:, :, 1, :, :], expected_var1) diff --git a/tests/test_validate_aggregator.py b/tests/test_validate_aggregator.py index 4b0237a7..67f0aa68 100644 --- a/tests/test_validate_aggregator.py +++ b/tests/test_validate_aggregator.py @@ -10,9 +10,9 @@ from samudra.aggregator.validate.map import MapAggregator from samudra.aggregator.validate.snapshot import SnapshotAggregator from samudra.aggregator.validate.sub_aggregator import ValidateSubAggregator -from samudra.constants import TensorMap -from samudra.utils.ctx import GridContext -from samudra.utils.data import DataSource, Normalize +from samudra.constants import DataLayout +from samudra.utils.ctx import BatchGrid +from samudra.utils.data import BatchPreprocessor, CanonicalSource from samudra.utils.output import ValBatchOutput from samudra.utils.wandb import Metrics @@ -21,13 +21,13 @@ def val_batch_of( h: int, w: int, *, - tensor_map: TensorMap, + data_layout: DataLayout, hist: int = 0, batch_size: int = 1, ) -> ValBatchOutput: - """Create a dummy Validation Batch loss / data from a DataSource.""" - n_prog_base = len(tensor_map.prognostic_var_names) - n_boundary_base = len(tensor_map.boundary_var_names) + """Create a dummy Validation Batch loss / data from a CanonicalSource.""" + n_prog_base = len(data_layout.prognostic_var_names) + n_boundary_base = len(data_layout.boundary_var_names) n_prog = (hist + 1) * n_prog_base n_boundary = (hist + 1) * n_boundary_base @@ -40,7 +40,7 @@ def val_batch_of( input_data=torch.randn(batch_size, n_prog + n_boundary, h, w), target_data=torch.randn(batch_size, n_prog, h, w), gen_data=torch.randn(batch_size, n_prog, h, w), - ctx=GridContext( + ctx=BatchGrid( label_mask=torch.ones(n_prog, h, w), input_resolution_cpu=( torch.linspace(-90, 90, steps=h), @@ -55,15 +55,17 @@ def val_batch_of( return batch -def tensor_map_for(src: DataSource) -> TensorMap: - return TensorMap(dataset_spec=src.dataset_spec) +def data_layout_for(source: CanonicalSource) -> DataLayout: + return source.data_layout -def normalize_for(src: DataSource, tensor_map: TensorMap) -> Normalize: - return Normalize( - src, - tensor_map.prognostic_var_names, - tensor_map.boundary_var_names, +def preprocessor_for( + source: CanonicalSource, data_layout: DataLayout +) -> BatchPreprocessor: + return BatchPreprocessor( + source, + data_layout.prognostic_var_names, + data_layout.boundary_var_names, ) @@ -88,22 +90,24 @@ def record_batch( self.num_recordings += 1 -def test_val_aggregator__no_op__is_same_as_train_aggregator(dummy_src: DataSource): - tensor_map = tensor_map_for(dummy_src) - normalize = normalize_for(dummy_src, tensor_map) - val_batch = val_batch_of(*dummy_src.grid_size, tensor_map=tensor_map) +def test_val_aggregator__no_op__is_same_as_train_aggregator( + dummy_source: CanonicalSource, +): + data_layout = data_layout_for(dummy_source) + preprocessor = preprocessor_for(dummy_source, data_layout) + val_batch = val_batch_of(*dummy_source.grid_size, data_layout=data_layout) num_prog_channels = val_batch.loss_per_channel.shape[0] val_agg = ValidateAggregator( {}, hist=0, num_prognostic_channels=num_prog_channels, - tensor_map=tensor_map, - normalize=normalize, + data_layout=data_layout, + preprocessor=preprocessor, ) val_agg.record_validation_batch(val_batch) val_agg.record_validation_batch(val_batch) - train_agg = TrainAggregator(tensor_map) + train_agg = TrainAggregator(data_layout) train_agg.record_batch(val_batch) train_agg.record_batch(val_batch) @@ -114,23 +118,23 @@ def test_val_aggregator__no_op__is_same_as_train_aggregator(dummy_src: DataSourc def test_train_val_aggregator__with_fake_subagg__is_added_to_logs( - dummy_src: DataSource, + dummy_source: CanonicalSource, ): - tensor_map = tensor_map_for(dummy_src) - normalize = normalize_for(dummy_src, tensor_map) - val_batch = val_batch_of(*dummy_src.grid_size, tensor_map=tensor_map) + data_layout = data_layout_for(dummy_source) + preprocessor = preprocessor_for(dummy_source, data_layout) + val_batch = val_batch_of(*dummy_source.grid_size, data_layout=data_layout) num_prog_channels = val_batch.loss_per_channel.shape[0] val_agg = ValidateAggregator( {"fake": FakeSubAggregator()}, hist=0, num_prognostic_channels=num_prog_channels, - tensor_map=tensor_map, - normalize=normalize, + data_layout=data_layout, + preprocessor=preprocessor, ) val_agg.record_validation_batch(val_batch) val_agg.record_validation_batch(val_batch) - train_agg = TrainAggregator(tensor_map) + train_agg = TrainAggregator(data_layout) train_agg.record_batch(val_batch) train_agg.record_batch(val_batch) @@ -149,20 +153,20 @@ def test_train_val_aggregator__with_fake_subagg__is_added_to_logs( def test_val_aggregator__hist_gt_0__does_not_require_wetmask_target_shape_match( - dummy_src: DataSource, + dummy_source: CanonicalSource, ): - tensor_map = tensor_map_for(dummy_src) - normalize = normalize_for(dummy_src, tensor_map) + data_layout = data_layout_for(dummy_source) + preprocessor = preprocessor_for(dummy_source, data_layout) val_batch = val_batch_of( - *dummy_src.grid_size, tensor_map=tensor_map, hist=1, batch_size=2 + *dummy_source.grid_size, data_layout=data_layout, hist=1, batch_size=2 ) num_prog_channels = val_batch.loss_per_channel.shape[0] val_agg = ValidateAggregator( {"fake": FakeSubAggregator()}, hist=1, num_prognostic_channels=num_prog_channels, - tensor_map=tensor_map, - normalize=normalize, + data_layout=data_layout, + preprocessor=preprocessor, ) val_agg.record_validation_batch(val_batch) val_logs = val_agg.get_logs(label="test") @@ -170,19 +174,19 @@ def test_val_aggregator__hist_gt_0__does_not_require_wetmask_target_shape_match( def test_validation_aggregator__reduced_only__omits_image_logs( - dummy_src: DataSource, + dummy_source: CanonicalSource, ): - tensor_map = tensor_map_for(dummy_src) - normalize = normalize_for(dummy_src, tensor_map) - val_batch = val_batch_of(*dummy_src.grid_size, tensor_map=tensor_map) + data_layout = data_layout_for(dummy_source) + preprocessor = preprocessor_for(dummy_source, data_layout) + val_batch = val_batch_of(*dummy_source.grid_size, data_layout=data_layout) num_prog_channels = val_batch.loss_per_channel.shape[0] val_agg = Aggregator.get_validation_aggregator( - dummy_src.metadata, + dummy_source.metadata, hist=0, - area_weights=dummy_src.spherical_area_weights, + area_weights=dummy_source.spherical_area_weights, num_prognostic_channels=num_prog_channels, - tensor_map=tensor_map, - normalize=normalize, + data_layout=data_layout, + preprocessor=preprocessor, include_image_aggregators=False, ) @@ -195,11 +199,11 @@ def test_validation_aggregator__reduced_only__omits_image_logs( def test_snapshot_aggregator__non_main_rank__skips_plot_rendering( - dummy_src: DataSource, monkeypatch: pytest.MonkeyPatch + dummy_source: CanonicalSource, monkeypatch: pytest.MonkeyPatch ): - tensor_map = tensor_map_for(dummy_src) - val_batch = val_batch_of(*dummy_src.grid_size, tensor_map=tensor_map) - aggregator = SnapshotAggregator(dummy_src.metadata, hist=0) + data_layout = data_layout_for(dummy_source) + val_batch = val_batch_of(*dummy_source.grid_size, data_layout=data_layout) + aggregator = SnapshotAggregator(dummy_source.metadata, hist=0) monkeypatch.setattr( "samudra.aggregator.validate.snapshot.is_main_process", @@ -220,9 +224,9 @@ def test_snapshot_aggregator__non_main_rank__skips_plot_rendering( def test_map_aggregator__non_main_rank__still_reduces_but_skips_plot_rendering( - dummy_src: DataSource, monkeypatch: pytest.MonkeyPatch + dummy_source: CanonicalSource, monkeypatch: pytest.MonkeyPatch ): - aggregator = MapAggregator(dummy_src.metadata, hist=0) + aggregator = MapAggregator(dummy_source.metadata, hist=0) reduce_calls: list[torch.Tensor] = [] monkeypatch.setattr( @@ -239,7 +243,7 @@ def fake_all_reduce_mean(tensor: torch.Tensor) -> torch.Tensor: fake_all_reduce_mean, ) - data = torch.ones(2, 1, *dummy_src.grid_size) + data = torch.ones(2, 1, *dummy_source.grid_size) aggregator.record_batch( loss=torch.tensor(1.0), target_data={"foo": data}, diff --git a/tests/test_wandb.py b/tests/test_wandb.py index a68d789e..25ef615c 100644 --- a/tests/test_wandb.py +++ b/tests/test_wandb.py @@ -28,7 +28,7 @@ def model_dump(self): return {} -class DummyDataContainer: +class DummyDataBundle: train_sources: list[Any] = [] @@ -48,7 +48,7 @@ def fail_load(*args, **kwargs): assert logger.setup_run( str(checkpoint_path), cast(Any, DummyConfig(tmp_path)), - cast(Any, DummyDataContainer()), + cast(Any, DummyDataBundle()), ) == (None, None) @@ -75,7 +75,7 @@ def fake_init(**kwargs): assert logger.setup_run( str(checkpoint_path), cast(Any, DummyConfig(tmp_path)), - cast(Any, DummyDataContainer()), + cast(Any, DummyDataBundle()), ) == ("run-123", "resume-me") assert init_kwargs["resume"] == "must" diff --git a/tests/test_writer.py b/tests/test_writer.py index 003b27bb..294f6927 100644 --- a/tests/test_writer.py +++ b/tests/test_writer.py @@ -9,29 +9,29 @@ import torch import xarray as xr -from samudra.constants import TensorMap, build_om4_spec -from samudra.utils.data import Normalize +from samudra.constants import build_om4_layout +from samudra.utils.data import BatchPreprocessor from samudra.utils.writer import ZarrWriter -from tests.conftest import TEST_FULL_DATASET_SPEC +from tests.conftest import TEST_FULL_DATA_LAYOUT -# write() never touches normalization (record_batch does), so these tests drive -# the writer directly with a buffer and omit a real Normalize. -_NO_NORMALIZE = cast(Normalize, None) +# write() never touches preprocessing (record_batch does), so these tests drive +# the writer directly with a buffer and omit a real preprocessor. +_NO_PREPROCESSOR = cast(BatchPreprocessor, None) def _source_coords(ny, nx): """Coords as `get_coords_dict` should be returned: 1D lat/lon dims, plus the grid metadata that survives `with_lat_lon_coords` (areacello, dz, lev, ocean_fraction).""" - spec = TEST_FULL_DATASET_SPEC - n_lev = spec.num_prognostic_depth_levels + data_layout = TEST_FULL_DATA_LAYOUT + n_lev = data_layout.num_prognostic_depth_levels lat = xr.DataArray(np.linspace(-89, 89, ny), dims="lat") lon = xr.DataArray(np.linspace(0, 359, nx), dims="lon") return { "lat": lat, "lon": lon, - "lev": xr.DataArray(list(spec.depth_levels), dims="lev"), - "dz": xr.DataArray(list(spec.depth_thickness), dims="lev"), + "lev": xr.DataArray(list(data_layout.depth_levels), dims="lev"), + "dz": xr.DataArray(list(data_layout.depth_thickness), dims="lev"), "areacello": xr.DataArray(np.ones((ny, nx)), dims=["lat", "lon"]), "ocean_fraction": xr.DataArray( np.ones((n_lev, ny, nx)), dims=["lev", "lat", "lon"] @@ -44,10 +44,9 @@ def _source_coords(ny, nx): def test_writer_output_is_analysis_ready(tmp_path): """The eval writer emits depth-stacked vars on y/x dims with grid metadata.""" - spec = TEST_FULL_DATASET_SPEC - tensor_map = TensorMap(dataset_spec=spec) - names = list(tensor_map.prognostic_var_names) - n_channels, n_lev = len(names), spec.num_prognostic_depth_levels + data_layout = TEST_FULL_DATA_LAYOUT + names = list(data_layout.prognostic_var_names) + n_channels, n_lev = len(names), data_layout.num_prognostic_depth_levels nt, ny, nx = 2, 3, 4 coords = _source_coords(ny, nx) @@ -57,8 +56,8 @@ def test_writer_output_is_analysis_ready(tmp_path): hist=0, model_path="dummy.ckpt", time_chunk_size=4, - normalize=_NO_NORMALIZE, - tensor_map=tensor_map, + preprocessor=_NO_PREPROCESSOR, + data_layout=data_layout, ) # buffer[t, c] is uniformly the channel index c, so each reassembled level can @@ -86,8 +85,8 @@ def test_writer_output_is_analysis_ready(tmp_path): assert (out["zos"].values == names.index("zos")).all() # (#3) grid metadata propagated, with horizontal dims renamed lat/lon -> y/x. - np.testing.assert_array_equal(out["lev"].values, spec.depth_levels) - np.testing.assert_array_equal(out["dz"].values, spec.depth_thickness) + np.testing.assert_array_equal(out["lev"].values, data_layout.depth_levels) + np.testing.assert_array_equal(out["dz"].values, data_layout.depth_thickness) assert out["areacello"].dims == ("y", "x") assert out["ocean_fraction"].dims == ("lev", "y", "x") # cell bounds propagate unchanged (enables dx/dy in analysis). @@ -108,9 +107,8 @@ def test_writer_prefers_real_2d_lat_lon(tmp_path): be rebuilt by broadcasting the 1D axes. `with_lat_lon_coords` preserves the real coords as `lat_2d`/`lon_2d`; the writer must emit those, not a broadcast. """ - spec = TEST_FULL_DATASET_SPEC - tensor_map = TensorMap(dataset_spec=spec) - n_channels = len(tensor_map.prognostic_var_names) + data_layout = TEST_FULL_DATA_LAYOUT + n_channels = len(data_layout.prognostic_var_names) ny, nx = 3, 4 coords = _source_coords(ny, nx) @@ -129,8 +127,8 @@ def test_writer_prefers_real_2d_lat_lon(tmp_path): hist=0, model_path="dummy.ckpt", time_chunk_size=4, - normalize=_NO_NORMALIZE, - tensor_map=tensor_map, + preprocessor=_NO_PREPROCESSOR, + data_layout=data_layout, ) writer.buffer = torch.zeros(1, n_channels, ny, nx) writer.time_buffer = xr.DataArray(np.arange(1), dims="time") @@ -148,9 +146,8 @@ def test_writer_prefers_real_2d_lat_lon(tmp_path): def test_writer_appends_along_time(tmp_path): """A second write extends the time axis without disturbing other coords.""" - spec = TEST_FULL_DATASET_SPEC - tensor_map = TensorMap(dataset_spec=spec) - n_channels = len(tensor_map.prognostic_var_names) + data_layout = TEST_FULL_DATA_LAYOUT + n_channels = len(data_layout.prognostic_var_names) ny, nx = 3, 4 coords = _source_coords(ny, nx) @@ -160,8 +157,8 @@ def test_writer_appends_along_time(tmp_path): hist=0, model_path="dummy.ckpt", time_chunk_size=4, - normalize=_NO_NORMALIZE, - tensor_map=tensor_map, + preprocessor=_NO_PREPROCESSOR, + data_layout=data_layout, ) def _write(times): @@ -185,20 +182,19 @@ def test_writer_shallow_spec_slices_depth_metadata(tmp_path): (e.g. thermo_dynamic_5) emits fewer levels. The writer slices the depth-resolved coords to the emitted level count instead of raising on a conflicting `lev` dim. """ - spec = build_om4_spec( + data_layout = build_om4_layout( prognostic_vars_key="thermo_dynamic_5", boundary_vars_key="tau_hfds" ) - tensor_map = TensorMap(dataset_spec=spec) - n_prog = spec.num_prognostic_depth_levels # 5, fewer than the source's 19 - n_channels = len(tensor_map.prognostic_var_names) + n_prog = data_layout.num_prognostic_depth_levels + n_channels = len(data_layout.prognostic_var_names) ny, nx, nt = 3, 4, 1 # Source coords carry the FULL 19-level depth metadata. coords = { "lat": xr.DataArray(np.linspace(-89, 89, ny), dims="lat"), "lon": xr.DataArray(np.linspace(0, 359, nx), dims="lon"), - "lev": xr.DataArray(list(spec.depth_levels), dims="lev"), - "dz": xr.DataArray(list(spec.depth_thickness), dims="lev"), + "lev": xr.DataArray(list(data_layout.depth_levels), dims="lev"), + "dz": xr.DataArray(list(data_layout.depth_thickness), dims="lev"), "ocean_fraction": xr.DataArray( np.ones((19, ny, nx)), dims=["lev", "lat", "lon"] ), @@ -210,8 +206,8 @@ def test_writer_shallow_spec_slices_depth_metadata(tmp_path): hist=0, model_path="dummy.ckpt", time_chunk_size=4, - normalize=_NO_NORMALIZE, - tensor_map=tensor_map, + preprocessor=_NO_PREPROCESSOR, + data_layout=data_layout, ) writer.buffer = torch.zeros(nt, n_channels, ny, nx) writer.time_buffer = xr.DataArray(np.arange(nt), dims="time") @@ -224,8 +220,10 @@ def test_writer_shallow_spec_slices_depth_metadata(tmp_path): assert out["thetao"].sizes["lev"] == n_prog assert out["dz"].sizes["lev"] == n_prog assert out["ocean_fraction"].sizes["lev"] == n_prog - np.testing.assert_array_equal(out["lev"].values, spec.depth_levels[:n_prog]) - np.testing.assert_array_equal(out["dz"].values, spec.depth_thickness[:n_prog]) + np.testing.assert_array_equal(out["lev"].values, data_layout.depth_levels[:n_prog]) + np.testing.assert_array_equal( + out["dz"].values, data_layout.depth_thickness[:n_prog] + ) def test_writer_curvilinear_grid_without_real_coords_raises(tmp_path): @@ -236,13 +234,12 @@ def test_writer_curvilinear_grid_without_real_coords_raises(tmp_path): coordinates. (The gaussian grid broadcasts fine -- see test_writer_output_is_analysis_ready, which drives the same coords.) """ - spec = build_om4_spec( + data_layout = build_om4_layout( prognostic_vars_key="thermo_dynamic_all", boundary_vars_key="tau_hfds_hfds_anom", grid_type="tripolar", ) - tensor_map = TensorMap(dataset_spec=spec) - n_channels = len(tensor_map.prognostic_var_names) + n_channels = len(data_layout.prognostic_var_names) ny, nx = 3, 4 coords = _source_coords(ny, nx) # 1D lat/lon only; no lat_2d/lon_2d @@ -252,8 +249,8 @@ def test_writer_curvilinear_grid_without_real_coords_raises(tmp_path): hist=0, model_path="dummy.ckpt", time_chunk_size=4, - normalize=_NO_NORMALIZE, - tensor_map=tensor_map, + preprocessor=_NO_PREPROCESSOR, + data_layout=data_layout, ) writer.buffer = torch.zeros(1, n_channels, ny, nx) writer.time_buffer = xr.DataArray(np.arange(1), dims="time")