(cfg)
| 22 | |
| 23 | |
| 24 | def build_diffusion(cfg): |
| 25 | beta_scheduler = cfg['beta_scheduler'] |
| 26 | diffusion_steps = cfg['diffusion_steps'] |
| 27 | |
| 28 | betas = get_named_beta_schedule(beta_scheduler, diffusion_steps) |
| 29 | model_mean_type = { |
| 30 | 'start_x': ModelMeanType.START_X, |
| 31 | 'previous_x': ModelMeanType.PREVIOUS_X, |
| 32 | 'epsilon': ModelMeanType.EPSILON |
| 33 | }[cfg['model_mean_type']] |
| 34 | model_var_type = { |
| 35 | 'learned': ModelVarType.LEARNED, |
| 36 | 'fixed_small': ModelVarType.FIXED_SMALL, |
| 37 | 'fixed_large': ModelVarType.FIXED_LARGE, |
| 38 | 'learned_range': ModelVarType.LEARNED_RANGE |
| 39 | }[cfg['model_var_type']] |
| 40 | if cfg.get('respace', None) is not None: |
| 41 | diffusion = SpacedDiffusion(use_timesteps=space_timesteps( |
| 42 | diffusion_steps, cfg['respace']), |
| 43 | betas=betas, |
| 44 | model_mean_type=model_mean_type, |
| 45 | model_var_type=model_var_type, |
| 46 | loss_type=LossType.MSE) |
| 47 | else: |
| 48 | diffusion = GaussianDiffusion(betas=betas, |
| 49 | model_mean_type=model_mean_type, |
| 50 | model_var_type=model_var_type, |
| 51 | loss_type=LossType.MSE) |
| 52 | return diffusion |
| 53 | |
| 54 | |
| 55 | @ARCHITECTURES.register_module() |
no test coverage detected