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