File size: 1,279 Bytes
d4cbafd
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
"""Eval MID SDD baseline checkpoint with ddim/stride=5 (20 denoising steps)."""
import argparse, yaml
from easydict import EasyDict
from mid import MID

if __name__ == '__main__':
    p = argparse.ArgumentParser()
    p.add_argument('--config', default='configs/baseline_sdd.yaml')
    p.add_argument('--dataset', default='sddd')
    p.add_argument('--eval_at', type=int, required=True)
    p.add_argument('--sampling', type=str, default='ddim')
    p.add_argument('--stride', type=int, default=5,
                   help='stride in reverse diffusion; #steps = 100/stride')
    args = p.parse_args()

    with open(args.config) as f:
        config = yaml.safe_load(f)
    config['config']   = args.config
    config['dataset']  = args.dataset
    config['exp_name'] = args.config.split('/')[-1].split('.')[0]
    config['dataset']  = args.dataset[:-1]
    config['eval_mode'] = True
    config['eval_at']   = args.eval_at
    cfg = EasyDict(config)

    agent = MID(cfg)
    # In diffusion.sample: stride = int(100/step), so pass step=100/stride
    step = 100 // args.stride    # stride=5 → step=20 → 20 DDIM steps
    print(f'Eval: sampling={args.sampling}, step arg={step} (stride={args.stride}, ~{100//args.stride} effective steps)')
    agent.eval(args.sampling, step)