Draws samples from an interpolated cosine timestep distribution (from simple diffusion).
(shape, image_d, noise_d_low, noise_d_high, sigma_data=1., min_value=1e-3, max_value=1e3, device='cpu', dtype=torch.float32)
| 182 | |
| 183 | |
| 184 | def rand_cosine_interpolated(shape, image_d, noise_d_low, noise_d_high, sigma_data=1., min_value=1e-3, max_value=1e3, device='cpu', dtype=torch.float32): |
| 185 | """Draws samples from an interpolated cosine timestep distribution (from simple diffusion).""" |
| 186 | |
| 187 | def logsnr_schedule_cosine(t, logsnr_min, logsnr_max): |
| 188 | t_min = math.atan(math.exp(-0.5 * logsnr_max)) |
| 189 | t_max = math.atan(math.exp(-0.5 * logsnr_min)) |
| 190 | return -2 * torch.log(torch.tan(t_min + t * (t_max - t_min))) |
| 191 | |
| 192 | def logsnr_schedule_cosine_shifted(t, image_d, noise_d, logsnr_min, logsnr_max): |
| 193 | shift = 2 * math.log(noise_d / image_d) |
| 194 | return logsnr_schedule_cosine(t, logsnr_min - shift, logsnr_max - shift) + shift |
| 195 | |
| 196 | def logsnr_schedule_cosine_interpolated(t, image_d, noise_d_low, noise_d_high, logsnr_min, logsnr_max): |
| 197 | logsnr_low = logsnr_schedule_cosine_shifted( |
| 198 | t, image_d, noise_d_low, logsnr_min, logsnr_max) |
| 199 | logsnr_high = logsnr_schedule_cosine_shifted( |
| 200 | t, image_d, noise_d_high, logsnr_min, logsnr_max) |
| 201 | return torch.lerp(logsnr_low, logsnr_high, t) |
| 202 | |
| 203 | logsnr_min = -2 * math.log(min_value / sigma_data) |
| 204 | logsnr_max = -2 * math.log(max_value / sigma_data) |
| 205 | u = stratified_uniform( |
| 206 | shape, group=0, groups=1, dtype=dtype, device=device |
| 207 | ) |
| 208 | logsnr = logsnr_schedule_cosine_interpolated( |
| 209 | u, image_d, noise_d_low, noise_d_high, logsnr_min, logsnr_max) |
| 210 | return torch.exp(-logsnr / 2) * sigma_data |
| 211 | |
| 212 | def rand_log_normal(shape, loc=0., scale=1., device='cpu', dtype=torch.float32): |
| 213 | """Draws samples from an lognormal distribution.""" |
no test coverage detected