| 113 | |
| 114 | |
| 115 | class _WrappedModel: |
| 116 | |
| 117 | def __init__(self, model, timestep_map, rescale_timesteps, |
| 118 | original_num_steps): |
| 119 | self.model = model |
| 120 | self.timestep_map = timestep_map |
| 121 | self.rescale_timesteps = rescale_timesteps |
| 122 | self.original_num_steps = original_num_steps |
| 123 | |
| 124 | def __call__(self, x, ts, **kwargs): |
| 125 | map_tensor = th.tensor(self.timestep_map, |
| 126 | device=ts.device, |
| 127 | dtype=ts.dtype) |
| 128 | new_ts = map_tensor[ts] |
| 129 | if self.rescale_timesteps: |
| 130 | # new_ts = new_ts.float() * (1000.0 / self.original_num_steps) |
| 131 | # new_ts = (new_ts * (100.0 / self.original_num_steps)).int() |
| 132 | |
| 133 | time_list = th.tensor([9,8,7,5,2], dtype=th.int32).to(new_ts.device) |
| 134 | new_ts = time_list[new_ts] |
| 135 | |
| 136 | return self.model(x, new_ts, **kwargs) |
| 137 | |
| 138 | def parameters(self): |
| 139 | return self.model.parameters() |