File size: 1,918 Bytes
bffb03e
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
"""Create small BCHW data with the multi-scale SatMAE++ protocol."""
import json
from pathlib import Path
import numpy as np
import torch
from torch.nn import functional as F
import yaml
import argparse

ROOT = Path(__file__).resolve().parents[1]

def main():
    parser = argparse.ArgumentParser(description=__doc__)
    parser.add_argument("--config", type=Path, default=ROOT / "conf/config.yaml")
    parser.add_argument("--output", type=Path, default=None)
    args = parser.parse_args()
    cfg = yaml.safe_load(args.config.read_text())
    d, rng = cfg["data"], np.random.default_rng(cfg["seed"])
    out = args.output or ROOT / d["root"]; out.mkdir(parents=True, exist_ok=True)
    shape = (d["train_samples"], d["channels"], d["image_size"], d["image_size"])
    test_shape = (d["test_samples"], d["channels"], d["image_size"], d["image_size"])
    def make_split(split_shape):
        images = rng.random(split_shape, dtype=np.float32)
        payload = {
            "images": images,
            "labels": rng.integers(d["num_classes"], size=split_shape[0], dtype=np.int64),
        }
        tensor = torch.from_numpy(images)
        for scale in d["scales"]:
            if scale != 1:
                payload[f"images_{scale}x"] = F.interpolate(
                    tensor, scale_factor=scale, mode="bilinear", align_corners=False
                ).numpy()
        return payload

    np.savez_compressed(out / "train.npz", **make_split(shape))
    np.savez_compressed(out / "test.npz", **make_split(test_shape))
    (out / "format.json").write_text(json.dumps({
        "format": "BCHW",
        "high_resolution_fields": [f"images_{scale}x" for scale in d["scales"] if scale != 1],
        "scales": d["scales"],
        "protocol": d["protocol"],
        "data_source": "synthetic",
    }, indent=2) + "\n")
    print("created", out / "train.npz", out / "test.npz")

if __name__ == "__main__": main()