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