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