(
model_output: torch.Tensor,
latents: torch.Tensor,
eta: float,
sigmas: torch.Tensor,
index: int,
prev_sample: torch.Tensor,
)
| 157 | |
| 158 | |
| 159 | def ddim_step( |
| 160 | model_output: torch.Tensor, |
| 161 | latents: torch.Tensor, |
| 162 | eta: float, |
| 163 | sigmas: torch.Tensor, |
| 164 | index: int, |
| 165 | prev_sample: torch.Tensor, |
| 166 | ): |
| 167 | model_output = convert_model_output(model_output, latents, sigmas, step_index=index) |
| 168 | prev_sample, prev_sample_mean, std_dev_t, dt_sqrt = ddim_update( |
| 169 | model_output, |
| 170 | sigmas.to(torch.float64), |
| 171 | index, |
| 172 | latents, |
| 173 | eta=eta, |
| 174 | ) |
| 175 | |
| 176 | # Compute log_prob |
| 177 | log_prob = ( |
| 178 | -((prev_sample.detach() - prev_sample_mean) ** 2) / (2 * ((std_dev_t * dt_sqrt) ** 2)) |
| 179 | - torch.log(std_dev_t * dt_sqrt) |
| 180 | - torch.log(torch.sqrt(2 * torch.as_tensor(math.pi))) |
| 181 | ) |
| 182 | |
| 183 | # mean along all but batch dimension |
| 184 | log_prob = log_prob.mean(dim=tuple(range(1, log_prob.ndim))) |
| 185 | return prev_sample, model_output, log_prob |
| 186 | |
| 187 | |
| 188 | @dataclass |
no test coverage detected