hakunamatata1997/audioEditing
2
1import torch2from tqdm import tqdm3# from torchvision import transforms as T4from typing import List, Optional, Dict, Union5from models import PipelineWrapper6 7 8def mu_tilde(model, xt, x0, timestep):9 "mu_tilde(x_t, x_0) DDPM paper eq. 7"10 prev_timestep = timestep - model.scheduler.config.num_train_timesteps // model.scheduler.num_inference_steps11 alpha_prod_t_prev = model.scheduler.alphas_cumprod[prev_timestep] if prev_timestep >= 0 \12 else model.scheduler.final_alpha_cumprod13 alpha_t = model.scheduler.alphas[timestep]14 beta_t = 1 - alpha_t15 alpha_bar = model.scheduler.alphas_cumprod[timestep]16 return ((alpha_prod_t_prev ** 0.5 * beta_t) / (1-alpha_bar)) * x0 + \17 ((alpha_t**0.5 * (1-alpha_prod_t_prev)) / (1 - alpha_bar)) * xt18 19 20def sample_xts_from_x0(model, x0, num_inference_steps=50, x_prev_mode=False):21 """22 Samples from P(x_1:T|x_0)23 """24 # torch.manual_seed(43256465436)25 alpha_bar = model.model.scheduler.alphas_cumprod26 sqrt_one_minus_alpha_bar = (1-alpha_bar) ** 0.527 alphas = model.model.scheduler.alphas28 # betas = 1 - alphas29 variance_noise_shape = (30 num_inference_steps + 1,31 model.model.unet.config.in_channels,32 # model.unet.sample_size,33 # model.unet.sample_size)34 x0.shape[-2],35 x0.shape[-1])36 37 timesteps = model.model.scheduler.timesteps.to(model.device)38 t_to_idx = {int(v): k for k, v in enumerate(timesteps)}39 xts = torch.zeros(variance_noise_shape).to(x0.device)40 xts[0] = x041 x_prev = x042 for t in reversed(timesteps):43 # idx = t_to_idx[int(t)]44 idx = num_inference_steps-t_to_idx[int(t)]45 if x_prev_mode:46 xts[idx] = x_prev * (alphas[t] ** 0.5) + torch.randn_like(x0) * ((1-alphas[t]) ** 0.5)47 x_prev = xts[idx].clone()48 else:49 xts[idx] = x0 * (alpha_bar[t] ** 0.5) + torch.randn_like(x0) * sqrt_one_minus_alpha_bar[t]50 # xts = torch.cat([xts, x0 ],dim = 0)51 52 return xts53 54 55def forward_step(model, model_output, timestep, sample):56 next_timestep = min(model.scheduler.config.num_train_timesteps - 2,57 timestep + model.scheduler.config.num_train_timesteps // model.scheduler.num_inference_steps)58 59 # 2. compute alphas, betas60 alpha_prod_t = model.scheduler.alphas_cumprod[timestep]61 # alpha_prod_t_next = self.scheduler.alphas_cumprod[next_timestep] if next_ltimestep >= 0 \62 # else self.scheduler.final_alpha_cumprod63 64 beta_prod_t = 1 - alpha_prod_t65 66 # 3. compute predicted original sample from predicted noise also called67 # "predicted x_0" of formula (12) from https://arxiv.org/pdf/2010.02502.pdf68 pred_original_sample = (sample - beta_prod_t ** (0.5) * model_output) / alpha_prod_t ** (0.5)69 70 # 5. TODO: simple noising implementatiom71 next_sample = model.scheduler.add_noise(pred_original_sample, model_output, torch.LongTensor([next_timestep]))72 return next_sample73 74 75def inversion_forward_process(model: PipelineWrapper,76 x0: torch.Tensor,77 etas: Optional[float] = None,78 prog_bar: bool = False,79 prompts: List[str] = [""],80 cfg_scales: List[float] = [3.5],81 num_inference_steps: int = 50,82 eps: Optional[float] = None,83 cutoff_points: Optional[List[float]] = None,84 numerical_fix: bool = False,85 extract_h_space: bool = False,86 extract_skipconns: bool = False,87 x_prev_mode: bool = False):88 if len(prompts) > 1 and extract_h_space:89 raise NotImplementedError("How do you split cfg_scales for hspace? TODO")90 91 if len(prompts) > 1 or prompts[0] != "":92 text_embeddings_hidden_states, text_embeddings_class_labels, \93 text_embeddings_boolean_prompt_mask = model.encode_text(prompts)94 # text_embeddings = encode_text(model, prompt)95 96 # # classifier free guidance97 batch_size = len(prompts)98 cfg_scales_tensor = torch.ones((batch_size, *x0.shape[1:]), device=model.device, dtype=x0.dtype)99 100 # if len(prompts) > 1:101 # if cutoff_points is None:102 # cutoff_points = [i * 1 / batch_size for i in range(1, batch_size)]103 # if len(cfg_scales) == 1:104 # cfg_scales *= batch_size105 # elif len(cfg_scales) < batch_size:106 # raise ValueError("Not enough target CFG scales")107 108 # cutoff_points = [int(x * cfg_scales_tensor.shape[2]) for x in cutoff_points]109 # cutoff_points = [0, *cutoff_points, cfg_scales_tensor.shape[2]]110 111 # for i, (start, end) in enumerate(zip(cutoff_points[:-1], cutoff_points[1:])):112 # cfg_scales_tensor[i, :, end:] = 0113 # cfg_scales_tensor[i, :, :start] = 0114 # cfg_scales_tensor[i] *= cfg_scales[i]115 # if prompts[i] == "":116 # cfg_scales_tensor[i] = 0117 # cfg_scales_tensor = T.functional.gaussian_blur(cfg_scales_tensor, kernel_size=15, sigma=1)118 # else:119 cfg_scales_tensor *= cfg_scales[0]120 121 uncond_embedding_hidden_states, uncond_embedding_class_lables, uncond_boolean_prompt_mask = model.encode_text([""])122 # uncond_embedding = encode_text(model, "")123 timesteps = model.model.scheduler.timesteps.to(model.device)124 variance_noise_shape = (125 num_inference_steps,126 model.model.unet.config.in_channels,127 # model.unet.sample_size,128 # model.unet.sample_size)129 x0.shape[-2],130 x0.shape[-1])131 132 if etas is None or (type(etas) in [int, float] and etas == 0):133 eta_is_zero = True134 zs = None135 else:136 eta_is_zero = False137 if type(etas) in [int, float]:138 etas = [etas]*model.model.scheduler.num_inference_steps139 xts = sample_xts_from_x0(model, x0, num_inference_steps=num_inference_steps, x_prev_mode=x_prev_mode)140 alpha_bar = model.model.scheduler.alphas_cumprod141 zs = torch.zeros(size=variance_noise_shape, device=model.device)142 hspaces = []143 skipconns = []144 t_to_idx = {int(v): k for k, v in enumerate(timesteps)}145 xt = x0146 # op = tqdm(reversed(timesteps)) if prog_bar else reversed(timesteps)147 op = tqdm(timesteps) if prog_bar else timesteps148 149 for t in op:150 # idx = t_to_idx[int(t)]151 idx = num_inference_steps - t_to_idx[int(t)] - 1152 # 1. predict noise residual153 if not eta_is_zero:154 xt = xts[idx+1][None]155 156 with torch.no_grad():157 out, out_hspace, out_skipconns = model.unet_forward(xt, timestep=t,158 encoder_hidden_states=uncond_embedding_hidden_states,159 class_labels=uncond_embedding_class_lables,160 encoder_attention_mask=uncond_boolean_prompt_mask)161 # out = model.unet.forward(xt, timestep= t, encoder_hidden_states=uncond_embedding)162 if len(prompts) > 1 or prompts[0] != "":163 cond_out, cond_out_hspace, cond_out_skipconns = model.unet_forward(164 xt.expand(len(prompts), -1, -1, -1), timestep=t,165 encoder_hidden_states=text_embeddings_hidden_states,166 class_labels=text_embeddings_class_labels,167 encoder_attention_mask=text_embeddings_boolean_prompt_mask)168 # cond_out = model.unet.forward(xt, timestep=t, encoder_hidden_states = text_embeddings)169 170 if len(prompts) > 1 or prompts[0] != "":171 # # classifier free guidance172 noise_pred = out.sample + \173 (cfg_scales_tensor * (cond_out.sample - out.sample.expand(batch_size, -1, -1, -1))174 ).sum(axis=0).unsqueeze(0)175 if extract_h_space or extract_skipconns:176 noise_h_space = out_hspace + cfg_scales[0] * (cond_out_hspace - out_hspace)177 if extract_skipconns:178 noise_skipconns = {k: [out_skipconns[k][j] + cfg_scales[0] *179 (cond_out_skipconns[k][j] - out_skipconns[k][j])180 for j in range(len(out_skipconns[k]))]181 for k in out_skipconns}182 else:183 noise_pred = out.sample184 if extract_h_space or extract_skipconns:185 noise_h_space = out_hspace186 if extract_skipconns:187 noise_skipconns = out_skipconns188 if extract_h_space or extract_skipconns:189 hspaces.append(noise_h_space)190 if extract_skipconns:191 skipconns.append(noise_skipconns)192 193 if eta_is_zero:194 # 2. compute more noisy image and set x_t -> x_t+1195 xt = forward_step(model.model, noise_pred, t, xt)196 else:197 # xtm1 = xts[idx+1][None]198 xtm1 = xts[idx][None]199 # pred of x0200 if model.model.scheduler.config.prediction_type == 'epsilon':201 pred_original_sample = (xt - (1 - alpha_bar[t]) ** 0.5 * noise_pred) / alpha_bar[t] ** 0.5202 elif model.model.scheduler.config.prediction_type == 'v_prediction':203 pred_original_sample = (alpha_bar[t] ** 0.5) * xt - ((1 - alpha_bar[t]) ** 0.5) * noise_pred204 205 # direction to xt206 prev_timestep = t - model.model.scheduler.config.num_train_timesteps // \207 model.model.scheduler.num_inference_steps208 209 alpha_prod_t_prev = model.get_alpha_prod_t_prev(prev_timestep)210 variance = model.get_variance(t, prev_timestep)211 212 if model.model.scheduler.config.prediction_type == 'epsilon':213 radom_noise_pred = noise_pred214 elif model.model.scheduler.config.prediction_type == 'v_prediction':215 radom_noise_pred = (alpha_bar[t] ** 0.5) * noise_pred + ((1 - alpha_bar[t]) ** 0.5) * xt216 217 pred_sample_direction = (1 - alpha_prod_t_prev - etas[idx] * variance) ** (0.5) * radom_noise_pred218 219 mu_xt = alpha_prod_t_prev ** (0.5) * pred_original_sample + pred_sample_direction220 221 z = (xtm1 - mu_xt) / (etas[idx] * variance ** 0.5)222 223 zs[idx] = z224 225 # correction to avoid error accumulation226 if numerical_fix:227 xtm1 = mu_xt + (etas[idx] * variance ** 0.5)*z228 xts[idx] = xtm1229 230 if zs is not None:231 # zs[-1] = torch.zeros_like(zs[-1])232 zs[0] = torch.zeros_like(zs[0])233 # zs_cycle[0] = torch.zeros_like(zs[0])234 235 if extract_h_space:236 hspaces = torch.concat(hspaces, axis=0)237 return xt, zs, xts, hspaces238 239 if extract_skipconns:240 hspaces = torch.concat(hspaces, axis=0)241 return xt, zs, xts, hspaces, skipconns242 243 return xt, zs, xts244 245 246def reverse_step(model, model_output, timestep, sample, eta=0, variance_noise=None):247 # 1. get previous step value (=t-1)248 prev_timestep = timestep - model.model.scheduler.config.num_train_timesteps // \249 model.model.scheduler.num_inference_steps250 # 2. compute alphas, betas251 alpha_prod_t = model.model.scheduler.alphas_cumprod[timestep]252 alpha_prod_t_prev = model.get_alpha_prod_t_prev(prev_timestep)253 beta_prod_t = 1 - alpha_prod_t254 # 3. compute predicted original sample from predicted noise also called255 # "predicted x_0" of formula (12) from https://arxiv.org/pdf/2010.02502.pdf256 if model.model.scheduler.config.prediction_type == 'epsilon':257 pred_original_sample = (sample - beta_prod_t ** (0.5) * model_output) / alpha_prod_t ** (0.5)258 elif model.model.scheduler.config.prediction_type == 'v_prediction':259 pred_original_sample = (alpha_prod_t ** 0.5) * sample - (beta_prod_t ** 0.5) * model_output260 261 # 5. compute variance: "sigma_t(η)" -> see formula (16)262 # σ_t = sqrt((1 − α_t−1)/(1 − α_t)) * sqrt(1 − α_t/α_t−1)263 # variance = self.scheduler._get_variance(timestep, prev_timestep)264 variance = model.get_variance(timestep, prev_timestep)265 # std_dev_t = eta * variance ** (0.5)266 # Take care of asymetric reverse process (asyrp)267 if model.model.scheduler.config.prediction_type == 'epsilon':268 model_output_direction = model_output269 elif model.model.scheduler.config.prediction_type == 'v_prediction':270 model_output_direction = (alpha_prod_t**0.5) * model_output + (beta_prod_t**0.5) * sample271 # 6. compute "direction pointing to x_t" of formula (12) from https://arxiv.org/pdf/2010.02502.pdf272 # pred_sample_direction = (1 - alpha_prod_t_prev - std_dev_t**2) ** (0.5) * model_output_direction273 pred_sample_direction = (1 - alpha_prod_t_prev - eta * variance) ** (0.5) * model_output_direction274 # 7. compute x_t without "random noise" of formula (12) from https://arxiv.org/pdf/2010.02502.pdf275 prev_sample = alpha_prod_t_prev ** (0.5) * pred_original_sample + pred_sample_direction276 # 8. Add noice if eta > 0277 if eta > 0:278 if variance_noise is None:279 variance_noise = torch.randn(model_output.shape, device=model.device)280 sigma_z = eta * variance ** (0.5) * variance_noise281 prev_sample = prev_sample + sigma_z282 283 return prev_sample284 285 286def inversion_reverse_process(model: PipelineWrapper,287 xT: torch.Tensor,288 skips: torch.Tensor,289 fix_alpha: float = 0.1,290 etas: float = 0,291 prompts: List[str] = [""],292 neg_prompts: List[str] = [""],293 cfg_scales: Optional[List[float]] = None,294 prog_bar: bool = False,295 zs: Optional[List[torch.Tensor]] = None,296 # controller=None,297 cutoff_points: Optional[List[float]] = None,298 hspace_add: Optional[torch.Tensor] = None,299 hspace_replace: Optional[torch.Tensor] = None,300 skipconns_replace: Optional[Dict[int, torch.Tensor]] = None,301 zero_out_resconns: Optional[Union[int, List]] = None,302 asyrp: bool = False,303 extract_h_space: bool = False,304 extract_skipconns: bool = False):305 306 batch_size = len(prompts)307 308 text_embeddings_hidden_states, text_embeddings_class_labels, \309 text_embeddings_boolean_prompt_mask = model.encode_text(prompts)310 uncond_embedding_hidden_states, uncond_embedding_class_lables, \311 uncond_boolean_prompt_mask = model.encode_text(neg_prompts)312 # text_embeddings = encode_text(model, prompts)313 # uncond_embedding = encode_text(model, [""] * batch_size)314 315 masks = torch.ones((batch_size, *xT.shape[1:]), device=model.device, dtype=xT.dtype)316 cfg_scales_tensor = torch.ones((batch_size, *xT.shape[1:]), device=model.device, dtype=xT.dtype)317 318 # if batch_size > 1:319 # if cutoff_points is None:320 # cutoff_points = [i * 1 / batch_size for i in range(1, batch_size)]321 # if len(cfg_scales) == 1:322 # cfg_scales *= batch_size323 # elif len(cfg_scales) < batch_size:324 # raise ValueError("Not enough target CFG scales")325 326 # cutoff_points = [int(x * cfg_scales_tensor.shape[2]) for x in cutoff_points]327 # cutoff_points = [0, *cutoff_points, cfg_scales_tensor.shape[2]]328 329 # for i, (start, end) in enumerate(zip(cutoff_points[:-1], cutoff_points[1:])):330 # cfg_scales_tensor[i, :, end:] = 0331 # cfg_scales_tensor[i, :, :start] = 0332 # masks[i, :, end:] = 0333 # masks[i, :, :start] = 0334 # cfg_scales_tensor[i] *= cfg_scales[i]335 # cfg_scales_tensor = T.functional.gaussian_blur(cfg_scales_tensor, kernel_size=15, sigma=1)336 # masks = T.functional.gaussian_blur(masks, kernel_size=15, sigma=1)337 # else:338 cfg_scales_tensor *= cfg_scales[0]339 340 if etas is None:341 etas = 0342 if type(etas) in [int, float]:343 etas = [etas]*model.model.scheduler.num_inference_steps344 assert len(etas) == model.model.scheduler.num_inference_steps345 timesteps = model.model.scheduler.timesteps.to(model.device)346 347 # xt = xT.expand(1, -1, -1, -1)348 xt = xT[skips.max()].unsqueeze(0)349 op = tqdm(timesteps[-zs.shape[0]:]) if prog_bar else timesteps[-zs.shape[0]:]350 351 t_to_idx = {int(v): k for k, v in enumerate(timesteps[-zs.shape[0]:])}352 hspaces = []353 skipconns = []354 355 for it, t in enumerate(op):356 # idx = t_to_idx[int(t)]357 idx = model.model.scheduler.num_inference_steps - t_to_idx[int(t)] - \358 (model.model.scheduler.num_inference_steps - zs.shape[0] + 1)359 # # Unconditional embedding360 with torch.no_grad():361 uncond_out, out_hspace, out_skipconns = model.unet_forward(362 xt, timestep=t,363 encoder_hidden_states=uncond_embedding_hidden_states,364 class_labels=uncond_embedding_class_lables,365 encoder_attention_mask=uncond_boolean_prompt_mask,366 mid_block_additional_residual=(None if hspace_add is None else367 (1 / (cfg_scales[0] + 1)) *368 (hspace_add[-zs.shape[0]:][it] if hspace_add.shape[0] > 1369 else hspace_add)),370 replace_h_space=(None if hspace_replace is None else371 (hspace_replace[-zs.shape[0]:][it].unsqueeze(0) if hspace_replace.shape[0] > 1372 else hspace_replace)),373 zero_out_resconns=zero_out_resconns,374 replace_skip_conns=(None if skipconns_replace is None else375 (skipconns_replace[-zs.shape[0]:][it] if len(skipconns_replace) > 1376 else skipconns_replace))377 ) # encoder_hidden_states = uncond_embedding)378 379 # # Conditional embedding380 if prompts:381 with torch.no_grad():382 cond_out, cond_out_hspace, cond_out_skipconns = model.unet_forward(383 xt.expand(batch_size, -1, -1, -1),384 timestep=t,385 encoder_hidden_states=text_embeddings_hidden_states,386 class_labels=text_embeddings_class_labels,387 encoder_attention_mask=text_embeddings_boolean_prompt_mask,388 mid_block_additional_residual=(None if hspace_add is None else389 (cfg_scales[0] / (cfg_scales[0] + 1)) *390 (hspace_add[-zs.shape[0]:][it] if hspace_add.shape[0] > 1391 else hspace_add)),392 replace_h_space=(None if hspace_replace is None else393 (hspace_replace[-zs.shape[0]:][it].unsqueeze(0) if hspace_replace.shape[0] > 1394 else hspace_replace)),395 zero_out_resconns=zero_out_resconns,396 replace_skip_conns=(None if skipconns_replace is None else397 (skipconns_replace[-zs.shape[0]:][it] if len(skipconns_replace) > 1398 else skipconns_replace))399 ) # encoder_hidden_states = text_embeddings)400 401 z = zs[idx] if zs is not None else None402 # print(f'idx: {idx}')403 # print(f't: {t}')404 z = z.unsqueeze(0)405 # z = z.expand(batch_size, -1, -1, -1)406 if prompts:407 # # classifier free guidance408 # noise_pred = uncond_out.sample + cfg_scales_tensor * (cond_out.sample - uncond_out.sample)409 noise_pred = uncond_out.sample + \410 (cfg_scales_tensor * (cond_out.sample - uncond_out.sample.expand(batch_size, -1, -1, -1))411 ).sum(axis=0).unsqueeze(0)412 if extract_h_space or extract_skipconns:413 noise_h_space = out_hspace + cfg_scales[0] * (cond_out_hspace - out_hspace)414 if extract_skipconns:415 noise_skipconns = {k: [out_skipconns[k][j] + cfg_scales[0] *416 (cond_out_skipconns[k][j] - out_skipconns[k][j])417 for j in range(len(out_skipconns[k]))]418 for k in out_skipconns}419 else:420 noise_pred = uncond_out.sample421 if extract_h_space or extract_skipconns:422 noise_h_space = out_hspace423 if extract_skipconns:424 noise_skipconns = out_skipconns425 426 if extract_h_space or extract_skipconns:427 hspaces.append(noise_h_space)428 if extract_skipconns:429 skipconns.append(noise_skipconns)430 431 # 2. compute less noisy image and set x_t -> x_t-1432 xt = reverse_step(model, noise_pred, t, xt, eta=etas[idx], variance_noise=z)433 # if controller is not None:434 # xt = controller.step_callback(xt)435 436 # "fix" xt437 apply_fix = ((skips.max() - skips) > it)438 if apply_fix.any():439 apply_fix = (apply_fix * fix_alpha).unsqueeze(1).unsqueeze(2).unsqueeze(3).to(xT.device)440 xt = (masks * (xt.expand(batch_size, -1, -1, -1) * (1 - apply_fix) +441 apply_fix * xT[skips.max() - it - 1].expand(batch_size, -1, -1, -1))442 ).sum(axis=0).unsqueeze(0)443 444 if extract_h_space:445 return xt, zs, torch.concat(hspaces, axis=0)446 447 if extract_skipconns:448 return xt, zs, torch.concat(hspaces, axis=0), skipconns449 450 return xt, zs451 