"""Create paired low/high resolution scenes with labels for kNN evaluation.""" import argparse, json from pathlib import Path import numpy as np, yaml from PIL import Image ROOT = Path(__file__).resolve().parents[1] def make_split(count, cfg, rng): size, target, channels, classes = cfg["input_size"], cfg["target_size"], cfg["channels"], cfg["num_classes"] labels = np.arange(count, dtype=np.int64) % classes; rng.shuffle(labels) gsd = rng.choice(np.asarray(cfg["gsd_values"], dtype=np.float32), count) y, x = np.mgrid[:target, :target].astype(np.float32); images = np.empty((count, channels, size, size), np.float32); targets = np.empty((count, channels, target, target), np.float32) for i, label in enumerate(labels): pattern = np.sin((x + label*2)*np.pi*(label+1)/target) + np.cos((y-label*2)*np.pi*(label+1)/target) pattern = (pattern-pattern.min())/(pattern.max()-pattern.min()) scene = np.stack([np.roll(pattern, label*c, axis=c%2) for c in range(channels)]) targets[i] = np.clip(scene + rng.normal(0, .02 + .01*gsd[i], scene.shape), 0, 1) images[i] = np.asarray([Image.fromarray((targets[i,c]*255).astype('uint8')).resize((size,size), Image.Resampling.BOX) for c in range(channels)], dtype=np.float32)/255 return images, targets, labels, gsd def main(): p=argparse.ArgumentParser(); p.add_argument("--config", default=str(ROOT/"conf/config.yaml")); a=p.parse_args(); cfg=yaml.safe_load(Path(a.config).read_text()); d=cfg["data"]; out=ROOT/d["root"]; out.mkdir(exist_ok=True); rng=np.random.default_rng(cfg["seed"]) for split,n in (("train",d["train_samples"]),("test",d["test_samples"])): images,targets,labels,gsd=make_split(n,d,rng); np.savez_compressed(out/f"{split}.npz", images=images, targets=targets, labels=labels, gsd=gsd) (out/"format.json").write_text(json.dumps({"images":"float32 BCHW input resolution","targets":"float32 BCHW target resolution","labels":"int64","gsd":"metres per pixel"},indent=2)+"\n"); print("created",out/"train.npz",out/"test.npz") if __name__ == "__main__": main()