(self, use_timesteps: Iterable[int], **kwargs)
| 958 | """ |
| 959 | |
| 960 | def __init__(self, use_timesteps: Iterable[int], **kwargs): |
| 961 | self.use_timesteps = set(use_timesteps) |
| 962 | self.timestep_map = [] |
| 963 | self.original_num_steps = len(kwargs["betas"]) |
| 964 | |
| 965 | base_diffusion = GaussianDiffusion(**kwargs) # pylint: disable=missing-kwoa |
| 966 | last_alpha_cumprod = 1.0 |
| 967 | new_betas = [] |
| 968 | for i, alpha_cumprod in enumerate(base_diffusion.alphas_cumprod): |
| 969 | if i in self.use_timesteps: |
| 970 | new_betas.append(1 - alpha_cumprod / last_alpha_cumprod) |
| 971 | last_alpha_cumprod = alpha_cumprod |
| 972 | self.timestep_map.append(i) |
| 973 | kwargs["betas"] = np.array(new_betas) |
| 974 | super().__init__(**kwargs) |
| 975 | |
| 976 | def p_mean_variance(self, model, *args, **kwargs): |
| 977 | return super().p_mean_variance(self._wrap_model(model), *args, **kwargs) |
nothing calls this directly
no test coverage detected