File size: 2,761 Bytes
80cf062 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 | from __future__ import annotations
import argparse
from pathlib import Path
import sys
import h5py
import numpy as np
import yaml
PROJECT_ROOT = Path(__file__).resolve().parents[1]
sys.path.insert(0, str(PROJECT_ROOT))
sys.path.insert(0, str(PROJECT_ROOT / "scripts"))
from era5_adapter import inspect_era5_contract
def load_config(path: Path) -> dict:
with path.open("r", encoding="utf-8") as handle:
return yaml.safe_load(handle)
def generate_year_file(
output_path: Path,
variables: list[str],
time_steps: int,
height: int,
width: int,
time_step_hours: int,
) -> None:
output_path.parent.mkdir(parents=True, exist_ok=True)
channels = len(variables)
means = np.zeros((1, channels, 1, 1), dtype=np.float32)
stds = np.ones((1, channels, 1, 1), dtype=np.float32)
with h5py.File(output_path, "w") as handle:
fields = handle.create_dataset(
"fields",
shape=(time_steps, channels, height, width),
dtype="float32",
chunks=(1, 1, height, width),
fillvalue=0.0,
)
fields.attrs["variables"] = variables
fields.attrs["time_step"] = time_step_hours
handle.create_dataset("global_means", data=means)
handle.create_dataset("global_stds", data=stds)
def main() -> None:
parser = argparse.ArgumentParser(description="Generate lightweight ERA5-style HDF5 files for W-MAE.")
parser.add_argument("--config", type=Path, default=PROJECT_ROOT / "conf" / "config.yaml")
parser.add_argument("--output-dir", type=Path, default=None)
args = parser.parse_args()
config = load_config(args.config)
data_config = config["data"]
fake_config = config["fake_data"]
output_dir = args.output_dir or PROJECT_ROOT / data_config["dataset_dir"]
variables = list(fake_config["variables"])
years = list(fake_config["years"])
source_height, source_width = data_config["source_size"]
for year in years:
path = output_dir / "data" / f"{year}.h5"
generate_year_file(
output_path=path,
variables=variables,
time_steps=int(fake_config["time_steps"]),
height=int(source_height),
width=int(source_width),
time_step_hours=int(data_config["time_step_hours"]),
)
print(f"created {path} ({path.stat().st_size / 1024:.1f} KiB physical size)")
report = inspect_era5_contract(output_dir, years, variables)
print(f"validated fields shape: {report['fields_shape']}")
print(f"W-MAE spatial adapter: {report['crop']} -> {report['model_size']}")
print("Synthetic channel names are test-only and must not be used for real training.")
if __name__ == "__main__":
main()
|