sra-trajectory-code / MID /eval_sdd_baseline.py
po03087's picture
SRA: MID/LED/MoFlow code + RUNNING.md instructions (code only, no data/ckpts)
d4cbafd verified
Raw
History Blame Contribute Delete
1.28 kB
"""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)