(
self,
features: torch.Tensor,
schedule=[1.0, 0.5],
generator=None,
)
| 127 | |
| 128 | @torch.no_grad() |
| 129 | def __call__( |
| 130 | self, |
| 131 | features: torch.Tensor, |
| 132 | schedule=[1.0, 0.5], |
| 133 | generator=None, |
| 134 | ): |
| 135 | features = self.ldm_transform_latent(features) |
| 136 | ts = self.round_timesteps( |
| 137 | torch.arange(0, 1024), |
| 138 | 1024, |
| 139 | self.n_distilled_steps, |
| 140 | truncate_start=False, |
| 141 | ) |
| 142 | shape = ( |
| 143 | features.size(0), |
| 144 | 3, |
| 145 | 8 * features.size(2), |
| 146 | 8 * features.size(3), |
| 147 | ) |
| 148 | x_start = torch.zeros(shape, device=features.device, dtype=features.dtype) |
| 149 | schedule_timesteps = [int((1024 - 1) * s) for s in schedule] |
| 150 | for i in schedule_timesteps: |
| 151 | t = ts[i].item() |
| 152 | t_ = torch.tensor([t] * features.shape[0]).to(self.device) |
| 153 | # noise = torch.randn_like(x_start) |
| 154 | noise = torch.randn(x_start.shape, dtype=x_start.dtype, generator=generator).to(device=x_start.device) |
| 155 | x_start = ( |
| 156 | _extract_into_tensor(self.sqrt_alphas_cumprod, t_, x_start.shape) * x_start |
| 157 | + _extract_into_tensor(self.sqrt_one_minus_alphas_cumprod, t_, x_start.shape) * noise |
| 158 | ) |
| 159 | c_in = _extract_into_tensor(self.c_in, t_, x_start.shape) |
| 160 | |
| 161 | import torch.nn.functional as F |
| 162 | |
| 163 | from diffusers import UNet2DModel |
| 164 | |
| 165 | if isinstance(self.ckpt, UNet2DModel): |
| 166 | input = torch.concat([c_in * x_start, F.upsample_nearest(features, scale_factor=8)], dim=1) |
| 167 | model_output = self.ckpt(input, t_).sample |
| 168 | else: |
| 169 | model_output = self.ckpt(c_in * x_start, t_, features=features) |
| 170 | |
| 171 | B, C = x_start.shape[:2] |
| 172 | model_output, _ = torch.split(model_output, C, dim=1) |
| 173 | pred_xstart = ( |
| 174 | _extract_into_tensor(self.c_out, t_, x_start.shape) * model_output |
| 175 | + _extract_into_tensor(self.c_skip, t_, x_start.shape) * x_start |
| 176 | ).clamp(-1, 1) |
| 177 | x_start = pred_xstart |
| 178 | return x_start |
| 179 | |
| 180 | |
| 181 | def save_image(image, name): |
nothing calls this directly
no test coverage detected