| """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) |
| |
| step = 100 // args.stride |
| print(f'Eval: sampling={args.sampling}, step arg={step} (stride={args.stride}, ~{100//args.stride} effective steps)') |
| agent.eval(args.sampling, step) |
|
|