from __future__ import annotations from pathlib import Path from typing import Any, Sequence import h5py EXPECTED_CHANNELS = 20 SOURCE_SIZE = (721, 1440) MODEL_SIZE = (720, 1440) def _decode_variables(values: Sequence[Any]) -> list[str]: return [value.decode() if isinstance(value, bytes) else str(value) for value in values] def validate_variables(variables: Sequence[str], expected_channels: int = EXPECTED_CHANNELS) -> list[str]: names = list(variables) if len(names) != expected_channels: raise ValueError( f"W-MAE requires exactly {expected_channels} explicitly ordered variables; " f"received {len(names)}. The official repository does not publish a complete " "channel-name mapping, so no default mapping is assumed." ) if len(set(names)) != len(names): raise ValueError("W-MAE variable names must be unique and explicitly ordered.") return names def inspect_era5_contract( dataset_dir: str | Path, years: Sequence[int], variables: Sequence[str], expected_channels: int = EXPECTED_CHANNELS, ) -> dict[str, Any]: """Validate the HDF5 metadata before importing the torch-based datapipe.""" dataset_dir = Path(dataset_dir) names = validate_variables(variables, expected_channels) missing_years = [year for year in years if not (dataset_dir / "data" / f"{year}.h5").is_file()] if missing_years: raise ValueError(f"ERA5 year files are missing: {missing_years}") first_file = dataset_dir / "data" / f"{years[0]}.h5" with h5py.File(first_file, "r") as handle: if "fields" not in handle: raise ValueError(f"{first_file} does not contain a 'fields' dataset.") fields = handle["fields"] if fields.ndim != 4: raise ValueError(f"fields must have shape [T,C,H,W], got {fields.shape}.") if tuple(fields.shape[-2:]) != SOURCE_SIZE: raise ValueError(f"W-MAE expects source grid {SOURCE_SIZE}, got {fields.shape[-2:]}.") if "variables" not in fields.attrs or "time_step" not in fields.attrs: raise ValueError("fields must define 'variables' and 'time_step' attributes.") available = _decode_variables(fields.attrs["variables"]) fields_shape = tuple(fields.shape) time_step = int(fields.attrs["time_step"]) missing_variables = [name for name in names if name not in available] if missing_variables: raise ValueError(f"Configured variables are absent from HDF5 metadata: {missing_variables}") if "global_means" not in handle or "global_stds" not in handle: stats_dir = dataset_dir / "stats" if not (stats_dir / "global_means.npy").is_file() or not (stats_dir / "global_stds.npy").is_file(): raise ValueError("Normalization statistics are missing from HDF5 and dataset_dir/stats.") return { "file": str(first_file), "fields_shape": fields_shape, "time_step": time_step, "selected_variables": names, "channel_indices": [available.index(name) for name in names], "crop": "fields[..., :720, :]", "model_size": MODEL_SIZE, } def _crop_last_latitude(value: Any) -> Any: if tuple(value.shape[-2:]) == MODEL_SIZE: return value if tuple(value.shape[-2:]) != SOURCE_SIZE: raise ValueError(f"Expected trailing spatial shape {SOURCE_SIZE}, got {tuple(value.shape[-2:])}.") return value[..., : MODEL_SIZE[0], :] class WMAEERA5Dataset: """Thin W-MAE adapter around OneScience ERA5Dataset. The official W-MAE loader removes the final latitude row from a 721x1440 ERA5 field. This wrapper preserves OneScience loading and normalization while applying that exact spatial convention. """ def __init__( self, dataset_dir: str | Path, years: Sequence[int], variables: Sequence[str], task: str = "forecast", input_steps: int = 1, output_steps: int = 1, normalize: bool = True, ) -> None: if task not in {"pretrain", "forecast"}: raise ValueError("task must be either 'pretrain' or 'forecast'.") names = validate_variables(variables) inspect_era5_contract(dataset_dir, years, names) try: from onescience.datapipes.climate.era5 import ERA5Dataset except (ImportError, OSError) as error: raise RuntimeError( "OneScience ERA5Dataset could not be imported. Verify the active " "OneScience/PyTorch runtime before constructing WMAEERA5Dataset." ) from error self.task = task self.dataset = ERA5Dataset( dataset_dir=str(dataset_dir), used_years=list(years), used_variables=names, input_steps=input_steps, output_steps=output_steps, normalize=normalize, ) def __len__(self) -> int: return len(self.dataset) def __getitem__(self, index: int) -> tuple[Any, Any, Any, int, list[str]]: invar, outvar, cos_zenith, step_idx, time_index = self.dataset[index] invar = _crop_last_latitude(invar) outvar = _crop_last_latitude(outvar) cos_zenith = _crop_last_latitude(cos_zenith) target = invar if self.task == "pretrain" else outvar return invar, target, cos_zenith, step_idx, time_index