File size: 2,014 Bytes
929e312
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
"""Generate virtual data in the native Samudra NPZ layout."""

try:
    from ._bootstrap import ROOT
except ImportError:
    from _bootstrap import ROOT

import argparse
from pathlib import Path

import numpy as np
import yaml

STATE_CHANNELS = 77
BOUNDARY_CHANNELS = 4


def generate(path: str | Path, time: int, height: int, width: int, seed: int) -> None:
    if time < 10:
        raise ValueError("time must be at least 10 for four-pass recurrent training")
    rng = np.random.default_rng(seed)
    prognostic = rng.standard_normal((time, STATE_CHANNELS, height, width), dtype=np.float32)
    boundary = rng.standard_normal((time, BOUNDARY_CHANNELS, height, width), dtype=np.float32)
    output = Path(path)
    output.parent.mkdir(parents=True, exist_ok=True)
    np.savez_compressed(output, prognostic=prognostic, boundary=boundary)
    print(f"saved native data: {output} prognostic={prognostic.shape} boundary={boundary.shape}")


def main() -> None:
    parser = argparse.ArgumentParser()
    parser.add_argument("--config", default="./conf/config.yaml")
    parser.add_argument("--train-output", default="./data/train.npz")
    parser.add_argument("--test-output", default="./data/test.npz")
    parser.add_argument("--time", type=int, default=None)
    parser.add_argument("--height", type=int, default=None)
    parser.add_argument("--width", type=int, default=None)
    parser.add_argument("--seed", type=int, default=None)
    args = parser.parse_args()
    with open(args.config, encoding="utf-8") as handle:
        config = yaml.safe_load(handle)
    fake = config.get("fake_data", {})
    time = args.time or int(fake.get("time", 12))
    height = args.height or int(fake.get("height", 32))
    width = args.width or int(fake.get("width", 64))
    seed = args.seed if args.seed is not None else int(config["project"].get("seed", 1))
    generate(args.train_output, time, height, width, seed)
    generate(args.test_output, time, height, width, seed + 1)


if __name__ == "__main__":
    main()