| 346 | |
| 347 | |
| 348 | class _WrappedModel: |
| 349 | def __init__(self, model, timestep_map, rescale_timesteps, original_num_steps): |
| 350 | self.model = model |
| 351 | self.timestep_map = timestep_map |
| 352 | self.rescale_timesteps = rescale_timesteps |
| 353 | self.original_num_steps = original_num_steps |
| 354 | |
| 355 | def __call__(self, x, ts, **kwargs): |
| 356 | map_tensor = torch.tensor(self.timestep_map, device=ts.device, dtype=ts.dtype) |
| 357 | new_ts = map_tensor[ts] |
| 358 | if self.rescale_timesteps: |
| 359 | new_ts = new_ts.float() * (1000.0 / self.original_num_steps) |
| 360 | return self.model(x, new_ts, **kwargs) |
| 361 | |
| 362 | |
| 363 | @register_sampler(name='ddpm') |