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