BABE-2 / model /sampler.py
Vansh Chugh
initial deploy
6dea0da
Raw
History Blame Contribute Delete
16 kB
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()