Surya / scripts /fake_data.py
zhangrenchao's picture
Upload Surya model package
a13f4b9 verified
Raw
History Blame Contribute Delete
2.58 kB
"""Create synthetic 13-channel SDO-like solar sequences."""
import json
from pathlib import Path
import numpy as np
import yaml
ROOT = Path(__file__).resolve().parents[1]
def make_split(count, cfg, rng):
total = cfg["input_steps"] + cfg["forecast_steps"]
size, channels = cfg["image_size"], cfg["channels"]
y, x = np.mgrid[-1:1:complex(size), -1:1:complex(size)].astype(np.float32)
disk = (x * x + y * y <= 0.92 ** 2).astype(np.float32)
sequences = np.empty((count, total, channels, size, size), dtype=np.float32)
activity = np.empty((count, total), dtype=np.float32)
for sample in range(count):
phase = rng.uniform(0, 2 * np.pi)
amplitude = rng.uniform(0.25, 0.8)
for step in range(total):
center_x = 0.48 * np.sin(phase + step * 0.18)
center_y = 0.28 * np.cos(phase + step * 0.12)
region = np.exp(-((x - center_x) ** 2 + (y - center_y) ** 2) / 0.035)
activity[sample, step] = region.sum() * amplitude
for channel in range(channels):
corona = np.exp(-(x * x + y * y) * (1.5 + channel * 0.05))
texture = 0.08 * np.sin((channel + 1) * x * 3 + phase + step * 0.1)
sequences[sample, step, channel] = np.clip(
disk * (0.15 + 0.35 * corona + amplitude * region * (0.5 + channel / channels) + texture)
+ rng.normal(0, 0.01, (size, size)), 0, 1)
return sequences[:, :cfg["input_steps"]], sequences[:, cfg["input_steps"]:], activity[:, cfg["input_steps"]:]
def main():
cfg = yaml.safe_load((ROOT / "conf/config.yaml").read_text())
data, rng = cfg["data"], np.random.default_rng(cfg["seed"])
output = ROOT / data["root"]; output.mkdir(exist_ok=True)
for split, count in (("train", data["train_samples"]), ("test", data["test_samples"])):
inputs, targets, activity = make_split(count, data, rng)
np.savez_compressed(output / f"{split}.npz", inputs=inputs, targets=targets, activity=activity)
(output / "format.json").write_text(json.dumps({
"inputs": "float32 [N, 2, 13, H, W]", "targets": "float32 [N, rollout, 13, H, W]",
"activity": "float32 integrated synthetic active-region signal", "protocol": data["protocol"],
"channel_names": data["channel_names"], "normalization": "signum_log then channel affine",
"channel_mean": data["channel_mean"], "channel_std": data["channel_std"]}, indent=2) + "\n")
print("created", output / "train.npz", output / "test.npz")
if __name__ == "__main__":
main()