Spaces:
Running on Zero
Running on Zero
| from tqdm import tqdm | |
| import torch | |
| from . import blind_bwe_utils | |
| class LPFOperator(): | |
| """ | |
| Parametric degradation-filter model, fitted during sampling to match the | |
| denoised estimate's spectrum to the observed recording's spectrum. | |
| """ | |
| def __init__(self, args, device) -> None: | |
| self.args = args | |
| self.device = device | |
| self.fcmin = self.args.tester.blind_bwe.fcmin | |
| if self.args.tester.blind_bwe.fcmax == "nyquist": | |
| self.fcmax = self.args.exp.sample_rate // 2 | |
| else: | |
| self.fcmax = self.args.tester.blind_bwe.fcmax | |
| self.Amin = self.args.tester.blind_bwe.Amin | |
| self.Amax = self.args.tester.blind_bwe.Amax | |
| if self.args.tester.blind_bwe.optimization.last_slope_fixed: | |
| self.args.tester.blind_bwe.initial_conditions.A_p[-1] = self.Amin | |
| if self.args.tester.blind_bwe.optimization.first_slope_fixed: | |
| self.args.tester.blind_bwe.initial_conditions.A_m[-1] = self.Amax | |
| self.params_fref = torch.Tensor([args.tester.blind_bwe.initial_conditions.fref]).to(device) | |
| self.params_fc_p = torch.Tensor(self.args.tester.blind_bwe.initial_conditions.fc_p).to(device) | |
| assert (self.params_fc_p > self.params_fref).all(), "fc_p must be greater than fref" | |
| self.params_fc_p = torch.nn.Parameter(self.params_fc_p) | |
| self.params_fc_m = torch.Tensor(self.args.tester.blind_bwe.initial_conditions.fc_m).to(device) | |
| assert (self.params_fc_m < self.params_fref).all(), "fc_m must be smaller than fref" | |
| self.params_fc_m = torch.nn.Parameter(self.params_fc_m) | |
| self.params_A_p = torch.Tensor(self.args.tester.blind_bwe.initial_conditions.A_p).to(device) | |
| self.params_A_p = torch.nn.Parameter(self.params_A_p) | |
| self.params_A_m = torch.Tensor(self.args.tester.blind_bwe.initial_conditions.A_m).to(device) | |
| self.params_A_m = torch.nn.Parameter(self.params_A_m) | |
| self.freqs = torch.fft.rfftfreq(self.args.tester.blind_bwe.NFFT, d=1 / self.args.exp.sample_rate).to(self.device) | |
| self.params = [self.params_fref, self.params_fc_p, self.params_fc_m, self.params_A_p, self.params_A_m] | |
| self.optimizer = torch.optim.Adam(self.params, lr=self.args.tester.blind_bwe.lr_filter) | |
| self.tol = self.args.tester.blind_bwe.optimization.tol | |
| def assign_params(self, params): | |
| assert len(params[1]) == (len(params[3]) - 1) | |
| assert len(params[2]) == (len(params[4]) - 1) | |
| self.params_fref = params[0] | |
| self.params_fc_p = params[1] | |
| self.params_fc_m = params[2] | |
| self.params_A_p = params[3] | |
| self.params_A_m = params[4] | |
| self.params = [self.params_fref, self.params_fc_p, self.params_fc_m, self.params_A_p, self.params_A_m] | |
| def degradation(self, x): | |
| return self.apply_filter_fcA(x) | |
| def apply_filter_fcA(self, x): | |
| H = blind_bwe_utils.design_filter_3(self.params, self.freqs, block_low_freq=self.args.tester.blind_bwe.optimization.block_low_freq) | |
| return blind_bwe_utils.apply_filter(x, H, self.args.tester.blind_bwe.NFFT) | |
| def stop(self, prev_params): | |
| decision = False | |
| if (torch.abs(self.params[0] - prev_params[0]).mean() < self.tol[0]): | |
| if (torch.abs(self.params[1] - prev_params[1]).mean() < self.tol[0]): | |
| if (torch.abs(self.params[2] - prev_params[2]).mean() < self.tol[0]): | |
| if (torch.abs(self.params[3] - prev_params[3]).mean() < self.tol[1]): | |
| if (torch.abs(self.params[4] - prev_params[4]).mean() < self.tol[1]): | |
| decision = True | |
| return decision | |
| def collapse_regularization(self): | |
| dist = [] | |
| dist.append(self.params[1][0] - self.params[0][0]) | |
| for i in range(1, len(self.params[1])): | |
| dist.append(self.params[1][i] - self.params[1][i - 1]) | |
| dist.append(self.params[1][-1] - self.fcmax) | |
| dist.append(self.params[0][0] - self.params[2][0]) | |
| for i in range(1, len(self.params[2])): | |
| dist.append(self.params[2][i] - self.params[2][i - 1]) | |
| dist.append(self.fcmin - self.params[2][-1]) | |
| beta = self.args.tester.collapse_regularization.beta | |
| gamma = self.args.tester.collapse_regularization.gamma | |
| cost = [torch.exp(-beta * x.abs()**gamma) for x in dist] | |
| return torch.stack(cost).sum() | |
| def limit_params(self): | |
| for i in range(len(self.params)): | |
| self.params[i].detach_() | |
| self.params[0][0] = torch.clamp(self.params[0][0], min=self.fcmin, max=self.fcmax) | |
| if self.args.tester.blind_bwe.optimization.clamp_fc: | |
| self.params[1][0] = torch.clamp(self.params[1][0], min=self.params[0][0] + 1e-3, max=self.fcmax) | |
| for k in range(1, len(self.params[1])): | |
| self.params[1][k] = torch.clamp(self.params[1][k], min=self.params[1][k - 1] + 1e-3, max=self.fcmax) | |
| self.params[2][0] = torch.clamp(self.params[2][0], min=self.fcmin, max=self.params[0][0] - 1e-3) | |
| for k in range(1, len(self.params[2])): | |
| self.params[2][k] = torch.clamp(self.params[2][k], min=self.fcmin, max=self.params[2][k - 1] - 1e-3) | |
| assert (self.params[1] <= self.params[0][0]).any() == False, f"fc_p must be greater than fref: {self.params[1]}, {self.params[0][0]}" | |
| assert (self.params[2] >= self.params[0][0]).any() == False, f"fc_m must be smaller than fre: {self.params[2]}, {self.params[0][0]}" | |
| assert (self.params[2] <= self.freqs[1]).any() == False, f"fc_m must be greater than the minimum frequency: {self.params[2]}, {self.freqs[1]}" | |
| assert (self.params[1] >= self.freqs[-1]).any() == False, f"fc_p must be smaller than the maximum frequency: {self.params[1]}, {self.freqs[-1]}" | |
| if self.args.tester.blind_bwe.optimization.clamp_A: | |
| if self.args.tester.blind_bwe.optimization.only_negative_Ap: | |
| self.params[3][0] = torch.clamp(self.params[3][0], min=self.Amin, max=0) | |
| for k in range(1, len(self.params[3])): | |
| self.params[3][k] = torch.clamp(self.params[3][k], min=self.Amin, max=self.params[3][k - 1] - 1e-1) | |
| else: | |
| for k in range(len(self.params[3])): | |
| self.params[3][k] = torch.clamp(self.params[3][k], min=self.Amin, max=self.Amax) | |
| for k in range(len(self.params[4])): | |
| self.params[4][k] = torch.clamp(self.params[4][k], min=self.Amin, max=self.Amax) | |
| if self.args.tester.blind_bwe.optimization.last_slope_fixed: | |
| self.params[3][-1] = -self.args.tester.blind_bwe.Alim | |
| if self.args.tester.blind_bwe.optimization.first_slope_fixed: | |
| self.params[4][-1] = self.args.tester.blind_bwe.Alim | |
| def optimizer_func(self, Xden, Y): | |
| """ | |
| Xden: STFT of denoised estimate. Y: STFT of observations. | |
| """ | |
| H = blind_bwe_utils.design_filter_3(self.params, self.freqs, block_low_freq=self.args.tester.blind_bwe.optimization.block_low_freq) | |
| return blind_bwe_utils.apply_filter_and_norm_STFTmag_fweighted(Xden, Y, H, self.args.tester.posterior_sampling.freq_weighting_filter) | |
| class AR_LPFOperator(LPFOperator): | |
| def __init__(self, args, device): | |
| super().__init__(args, device) | |
| self.mask = None | |
| def degradation(self, x): | |
| return self.mask * x + (1 - self.mask) * self.apply_filter_fcA(x) | |
| class BlindSampler(): | |
| """ | |
| EDM sampler with reconstruction guidance, doing joint denoising and | |
| blind degradation-filter estimation. | |
| """ | |
| def __init__(self, model, diff_params, args): | |
| self.model = model | |
| self.diff_params = diff_params | |
| self.args = args | |
| if not self.args.tester.diff_params.same_as_training: | |
| self.update_diff_params() | |
| self.order = self.args.tester.order | |
| self.xi = self.args.tester.posterior_sampling.xi | |
| self.data_consistency = self.args.tester.posterior_sampling.data_consistency | |
| self.nb_steps = self.args.tester.T | |
| self.start_sigma = self.args.tester.posterior_sampling.start_sigma | |
| if self.start_sigma == "None": | |
| self.start_sigma = None | |
| self.operator = None | |
| def loss_fn_rec(x_hat, x): | |
| diff = x_hat - x | |
| return (diff**2).sum() / 2 | |
| self.rec_distance = lambda x_hat, x: loss_fn_rec(x_hat, x) | |
| def update_diff_params(self): | |
| self.diff_params.sigma_min = self.args.tester.diff_params.sigma_min | |
| self.diff_params.sigma_max = self.args.tester.diff_params.sigma_max | |
| self.diff_params.ro = self.args.tester.diff_params.ro | |
| self.diff_params.sigma_data = self.args.tester.diff_params.sigma_data | |
| self.diff_params.Schurn = self.args.tester.diff_params.Schurn | |
| self.diff_params.Stmin = self.args.tester.diff_params.Stmin | |
| self.diff_params.Stmax = self.args.tester.diff_params.Stmax | |
| self.diff_params.Snoise = self.args.tester.diff_params.Snoise | |
| def get_rec_grads(self, x_hat, y, x, t_i): | |
| """ | |
| Gradient of the reconstruction error (in the degraded-signal domain) | |
| with respect to the current diffusion latent. | |
| """ | |
| if self.args.tester.posterior_sampling.annealing_y.use: | |
| mode = self.args.tester.posterior_sampling.annealing_y.mode | |
| if mode == "same_as_x": | |
| y = y + torch.randn_like(y) * t_i | |
| elif mode == "same_as_x_limited": | |
| t_min = torch.Tensor([self.args.tester.posterior_sampling.annealing_y.sigma_min]).to(y.device) | |
| t_y = torch.max(t_i, t_min) | |
| y = y + torch.randn_like(y) * t_y | |
| elif mode == "fixed": | |
| t_min = torch.Tensor([self.args.tester.posterior_sampling.annealing_y.sigma_min]).to(y.device) | |
| y = y + torch.randn_like(y) * t_min | |
| norm = self.rec_distance(self.operator.degradation(x_hat), y) | |
| rec_grads = torch.autograd.grad(outputs=norm.sum(), inputs=x) | |
| rec_grads = rec_grads[0] | |
| normalization = self.args.tester.posterior_sampling.normalization | |
| if normalization == "grad_norm": | |
| normguide = torch.norm(rec_grads) / self.args.exp.audio_len**0.5 | |
| elif normalization == "loss_norm": | |
| normguide = norm / self.args.exp.audio_len**0.5 | |
| s = self.xi / (normguide + 1e-6) | |
| return s * rec_grads / t_i, norm | |
| def get_denoised_estimate(self, x, t_i): | |
| x_hat = self.diff_params.denoiser(x, self.model, t_i.unsqueeze(-1)) | |
| if self.args.tester.filter_out_cqt_DC_Nyq: | |
| x_hat = self.model.CQTransform.apply_hpf_DC(x_hat) | |
| return x_hat | |
| def denoised2score(self, x_d0, x, t): | |
| return (x_d0 - x) / t**2 | |
| def move_timestep(self, x, t, gamma, Snoise=1): | |
| t_hat = t + gamma * t | |
| epsilon = torch.randn(x.shape).to(x.device) * Snoise | |
| x_hat = x + ((t_hat**2 - t**2)**(1 / 2)) * epsilon | |
| return x_hat, t_hat | |
| def fit_params(self, denoised_estimate, y): | |
| Xden = blind_bwe_utils.apply_stft(denoised_estimate, self.args.tester.blind_bwe.NFFT) | |
| Y = blind_bwe_utils.apply_stft(y, self.args.tester.blind_bwe.NFFT) | |
| for i in range(self.args.tester.blind_bwe.optimization.max_iter): | |
| for j in range(len(self.operator.params)): | |
| self.operator.params[j].requires_grad = True | |
| self.operator.optimizer.zero_grad() | |
| rec_loss = self.operator.optimizer_func(Xden, Y) | |
| if self.args.tester.collapse_regularization.use: | |
| cost = self.operator.collapse_regularization() | |
| loss = rec_loss + self.args.tester.collapse_regularization.lambda_reg * cost | |
| else: | |
| loss = rec_loss | |
| loss.backward() | |
| torch.nn.utils.clip_grad_norm_(self.operator.params, self.args.tester.blind_bwe.optimization.grad_clip) | |
| self.operator.optimizer.step() | |
| self.operator.limit_params() | |
| if i > 0: | |
| if self.operator.stop(prev_params): | |
| break | |
| prev_params = [self.operator.params[k].clone().detach() for k in range(len(self.operator.params))] | |
| def step(self, x, t_i, t_i_1, gamma_i, blind=False, y=None): | |
| if self.args.tester.posterior_sampling.SNR_observations != "None": | |
| snr = 10**(self.args.tester.posterior_sampling.SNR_observations / 10) | |
| sigma2_s = torch.var(y, -1) | |
| sigma = torch.sqrt(sigma2_s / snr).unsqueeze(-1) | |
| y = y + sigma * torch.randn(y.shape).to(y.device) | |
| x_hat, t_hat = self.move_timestep(x, t_i, gamma_i, self.diff_params.Snoise) | |
| x_hat.requires_grad_(True) | |
| x_den = self.get_denoised_estimate(x_hat, t_hat) | |
| x_den_2 = x_den.clone().detach() | |
| if blind: | |
| self.fit_params(x_den_2, y) | |
| if self.args.tester.posterior_sampling.xi > 0 and y is not None: | |
| rec_grads, rec_loss = self.get_rec_grads(x_den, y, x_hat, t_hat) | |
| else: | |
| rec_loss = 0 | |
| rec_grads = 0 | |
| x_hat.detach_() | |
| uncond_score = self.denoised2score(x_den_2, x_hat, t_hat) | |
| score = uncond_score - rec_grads | |
| d = -t_hat * score | |
| h = t_i_1 - t_hat | |
| if t_i_1 != 0 and self.order == 2: | |
| t_prime = t_i_1 | |
| x_prime = x_hat + h * d | |
| x_prime.requires_grad_(True) | |
| x_den = self.get_denoised_estimate(x_prime, t_prime) | |
| x_den_2 = x_den.clone().detach() | |
| if blind: | |
| self.fit_params(x_den_2, y) | |
| if self.xi > 0 and y is not None: | |
| rec_grads, rec_loss = self.get_rec_grads(x_den, y, x_prime, t_prime) | |
| else: | |
| rec_loss = 0 | |
| rec_grads = 0 | |
| x_prime.detach_() | |
| uncond_score = self.denoised2score(x_den_2, x_prime, t_prime) | |
| score = uncond_score - rec_grads | |
| d_prime = -t_prime * score | |
| x = (x_hat + h * ((1 / 2) * d + (1 / 2) * d_prime)) | |
| elif self.order == 1: | |
| x = x_hat + h * d | |
| return x, x_den_2, rec_loss, score, rec_grads | |
| def predict_blind_bwe_AR(self, ylpf, y_masked, mask=None, x_init=None, progress_cb=None): | |
| self.operator = AR_LPFOperator(self.args, ylpf.device) | |
| self.operator.mask = mask | |
| y = mask * y_masked + (1 - mask) * ylpf | |
| y = y.unsqueeze(0) | |
| self.y = y | |
| return self.predict(shape=y.shape, device=y.device, blind=True, x_init=x_init, progress_cb=progress_cb) | |
| def predict_blind_bwe(self, y, x_init=None, progress_cb=None): | |
| self.operator = LPFOperator(self.args, y.device) | |
| self.y = y | |
| return self.predict(shape=y.shape, device=y.device, blind=True, x_init=x_init, progress_cb=progress_cb) | |
| def predict(self, shape, device, blind=False, x_init=None, progress_cb=None): | |
| if self.start_sigma is None: | |
| t = self.diff_params.create_schedule(self.nb_steps).to(device) | |
| x = self.diff_params.sample_prior(shape, t[0]).to(device) | |
| else: | |
| t = self.diff_params.create_schedule_from_initial_t(self.start_sigma, self.nb_steps).to(self.y.device) | |
| if x_init is not None: | |
| x = x_init.to(device) + self.diff_params.sample_prior(shape, t[0]).to(device) | |
| else: | |
| x = self.y + self.diff_params.sample_prior(shape, t[0]).to(device) | |
| gamma = self.diff_params.get_gamma(t).to(device) | |
| for i in tqdm(range(0, self.nb_steps, 1)): | |
| out = self.step(x, t[i], t[i + 1], gamma[i], blind=blind, y=self.y) | |
| x, x_den, rec_loss, score, lh_score = out | |
| if progress_cb is not None: | |
| progress_cb(i, self.nb_steps) | |
| if blind: | |
| return x.detach(), self.operator.params | |
| else: | |
| return x.detach() | |