SatMAE-PP / scripts /fake_data.py
zhangrenchao's picture
Upload SatMAE++ model package
bffb03e verified
Raw
History Blame Contribute Delete
1.92 kB
"""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()