Samudra / scripts /fake_data.py
yzt15806542928's picture
Upload folder using huggingface_hub
929e312 verified
Raw
History Blame Contribute Delete
2.01 kB
"""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()