| """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() |
|
|