File size: 1,976 Bytes
db75703 | 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 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 | #!/usr/bin/env python3
"""
8-D conditions → MA generate (64) → SwinIR SR (224) → frozen proxy → metrics.
Example:
python generate_and_eval.py \\
--cond-values "0.0784,144.40,310.56,551.57,-0.3088,0.2504,5.1633,2.2935" \\
--n-samples 4 \\
--outdir runs/demo
"""
from __future__ import annotations
import argparse
from pathlib import Path
from ma.cond import parse_cond_values
from ma.pipeline import run
ROOT = Path(__file__).resolve().parent
def main() -> None:
p = argparse.ArgumentParser(description="MA generate + measure + eval")
p.add_argument(
"--cond-values",
type=str,
required=True,
help="8 floats: redshift,flux_g,flux_r,flux_z,shape_e1,shape_e2,shape_r,sersic",
)
p.add_argument("--n-samples", type=int, default=4)
p.add_argument("--outdir", type=Path, default=ROOT / "runs" / "demo")
p.add_argument("--device", type=str, default="cuda")
p.add_argument("--rank", type=int, default=0)
p.add_argument("--base-seed", type=int, default=42)
p.add_argument("--mc-samples", type=int, default=10)
p.add_argument("--batch-size", type=int, default=4)
p.add_argument(
"--sampling-timesteps",
type=int,
default=50,
help="Fast strided steps (default 50). Use 0 for full 1000-step DDPM.",
)
p.add_argument("--skip-generate", action="store_true")
p.add_argument("--skip-sr", action="store_true")
args = p.parse_args()
cond = parse_cond_values(args.cond_values)
run(
cond,
outdir=Path(args.outdir),
n_samples=int(args.n_samples),
device=args.device,
rank=int(args.rank),
base_seed=int(args.base_seed),
mc_samples=int(args.mc_samples),
batch_size=int(args.batch_size),
sampling_timesteps=int(args.sampling_timesteps),
skip_generate=bool(args.skip_generate),
skip_sr=bool(args.skip_sr),
)
if __name__ == "__main__":
main()
|