#!/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()