OneScience-Group/SatMAE
030
1"""Create temporary fMoW-style temporal tensors and labels."""2import json3from pathlib import Path4import numpy as np5import yaml6 7ROOT = Path(__file__).resolve().parents[1]8 9def main():10 config = yaml.safe_load((ROOT / "conf/config.yaml").read_text())11 d = config["data"]12 out = ROOT / d["root"]13 out.mkdir(exist_ok=True)14 rng = np.random.default_rng(config["seed"])15 def make_split(samples):16 shape = (samples, d["frames"], d["channels"], d["image_size"], d["image_size"])17 images = rng.random(shape, dtype=np.float32)18 timestamps = np.stack(19 (20 rng.integers(0, 21, size=(samples, d["frames"])),21 rng.integers(0, 12, size=(samples, d["frames"])),22 rng.integers(0, 24, size=(samples, d["frames"])),23 ),24 axis=-1,25 ).astype(np.float32)26 order = np.argsort(timestamps[..., 0] * 12 * 24 + timestamps[..., 1] * 24 + timestamps[..., 2], axis=1)27 images = np.take_along_axis(images, order[:, :, None, None, None], axis=1)28 timestamps = np.take_along_axis(timestamps, order[..., None], axis=1)29 labels = rng.integers(d["num_classes"], size=samples, dtype=np.int64)30 return images, timestamps, labels31 32 train = make_split(d["train_samples"])33 test = make_split(d["test_samples"])34 np.savez_compressed(out / "train.npz", images=train[0], timestamps=train[1], labels=train[2])35 np.savez_compressed(out / "test.npz", images=test[0], timestamps=test[1], labels=test[2])36 (out / "format.json").write_text(json.dumps({37 "format": "BTCHW",38 "timestamp_format": "BT3: year_offset_2002, month_zero_based, hour",39 "source_protocol": d["protocol"],40 "data_source": "synthetic",41 }, indent=2) + "\n")42 print("created", out / "train.npz", out / "test.npz")43 44if __name__ == "__main__":45 main()46 