Download scripts/era5_adapter.py from OneScience-Group/W-MAE: direct link, hf CLI and curl.
- Browser
- Download file 5.44 kB
-
https://huggingface.co/OneScience-Group/W-MAE/resolve/main/scripts/era5_adapter.py
- Command line
-
hf download hf://OneScience-Group/W-MAE/scripts/era5_adapter.py
-
curl -L -o era5_adapter.py https://huggingface.co/OneScience-Group/W-MAE/resolve/main/scripts/era5_adapter.py
5.44 kB
| 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 | |