W-MAE / scripts /era5_adapter.py
yzt15806542928's picture
Upload folder using huggingface_hub
80cf062 verified
Raw
History Blame Contribute Delete
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