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