"""Flow-matching objective + ODE sampler for the deg-anchored restoration bridge. Noisy-start rectified flow: x0 = degraded + sigma*eps, x1 = clean, x_t=(1-t)x0+t x1, target velocity v = x1 - x0. Diversity comes from eps; sampling integrates v from a degraded-anchored start -> a plausible clean (the restoration trajectory). """ from __future__ import annotations import torch import torch.nn.functional as F def fm_loss(model, deg, clean, stem_id, sigma: float): B = deg.shape[0] eps = torch.randn_like(clean) x0 = deg + sigma * eps t = torch.rand(B, device=deg.device) tt = t[:, None, None] z_t = (1 - tt) * x0 + tt * clean v_target = clean - x0 return F.mse_loss(model(z_t, t, deg, stem_id), v_target) @torch.inference_mode() def sample(model, deg, stem_id, steps: int, sigma: float, generator=None): eps = torch.randn(deg.shape, generator=generator, device=deg.device) z = deg + sigma * eps dt = 1.0 / steps for i in range(steps): t = torch.full((deg.shape[0],), i * dt, device=deg.device) z = z + dt * model(z, t, deg, stem_id) return z def fm_loss_mix(model, deg, clean, mix, stem_id, sigma: float): """Deg-anchored flow loss WITH mix conditioning (model takes z_t,t,deg,stem_id,mix).""" B = deg.shape[0] eps = torch.randn_like(clean) x0 = deg + sigma * eps t = torch.rand(B, device=deg.device) tt = t[:, None, None] z_t = (1 - tt) * x0 + tt * clean v_target = clean - x0 return F.mse_loss(model(z_t, t, deg, stem_id, mix), v_target) @torch.inference_mode() def sample_mix(model, deg, mix, stem_id, steps: int, sigma: float, generator=None): eps = torch.randn(deg.shape, generator=generator, device=deg.device) z = deg + sigma * eps dt = 1.0 / steps for i in range(steps): t = torch.full((deg.shape[0],), i * dt, device=deg.device) z = z + dt * model(z, t, deg, stem_id, mix) return z # ---- cheap latent metrics (validated proxies for decoded-audio distance) ---- def latent_cos_dist(pred, clean): # best proxy (within-clip Spearman ~0.94) p, c = pred.flatten(1), clean.flatten(1) return (1 - F.cosine_similarity(p, c, dim=1)).mean() def latent_relL2(pred, clean): p, c = pred.flatten(1), clean.flatten(1) return ((p - c).norm(dim=1) / (c.norm(dim=1) + 1e-9)).mean() def energy_ratio(pred, clean): p, c = pred.flatten(1), clean.flatten(1) return (p.norm(dim=1) / (c.norm(dim=1) + 1e-9)).mean()