CoolFace
Modelpublic

OneScience-Group/SatMAE

sourceHugging Facecc-by-nc-4.0updated 22d agoView on Hugging Face
0likes30downloads
fake_data.py46 linesDownload Raw Back to scripts
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