CoolFace
Modelpublic

OneScience-Group/Scale-MAE

sourceHugging Facecc-by-nc-4.0updated 21d agoView on Hugging Face
0likes26downloads
fake_data.py27 linesDownload Raw Back to scripts
1"""Create paired low/high resolution scenes with labels for kNN evaluation."""2import argparse, json3from pathlib import Path4import numpy as np, yaml5from PIL import Image6ROOT = Path(__file__).resolve().parents[1]7 8def make_split(count, cfg, rng):9    size, target, channels, classes = cfg["input_size"], cfg["target_size"], cfg["channels"], cfg["num_classes"]10    labels = np.arange(count, dtype=np.int64) % classes; rng.shuffle(labels)11    gsd = rng.choice(np.asarray(cfg["gsd_values"], dtype=np.float32), count)12    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)13    for i, label in enumerate(labels):14        pattern = np.sin((x + label*2)*np.pi*(label+1)/target) + np.cos((y-label*2)*np.pi*(label+1)/target)15        pattern = (pattern-pattern.min())/(pattern.max()-pattern.min())16        scene = np.stack([np.roll(pattern, label*c, axis=c%2) for c in range(channels)])17        targets[i] = np.clip(scene + rng.normal(0, .02 + .01*gsd[i], scene.shape), 0, 1)18        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)/25519    return images, targets, labels, gsd20 21def main():22    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"])23    for split,n in (("train",d["train_samples"]),("test",d["test_samples"])):24        images,targets,labels,gsd=make_split(n,d,rng); np.savez_compressed(out/f"{split}.npz", images=images, targets=targets, labels=labels, gsd=gsd)25    (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")26if __name__ == "__main__": main()27