VOSR / vosr.py
41-807's picture
Upload folder using huggingface_hub
93d565a verified
Raw
History Blame Contribute Delete
14.3 kB
import torch
import torch.nn as nn
from einops import rearrange
import numpy as np
import random
import torch.distributions as D
def _beta_ratio(size, alpha: float, beta: float, device="cpu"):
return D.Beta(alpha, beta).sample(size).to(device)
def stopgrad(x):
return x.detach().clone()
def _zero_like(x):
"""Create zero tensor(s) matching x, which can be Tensor or list of Tensors."""
if isinstance(x, (list, tuple)):
return [torch.zeros_like(xi) for xi in x]
return torch.zeros_like(x)
def _set_indices(x, indices, src):
"""x[indices] = src[indices] for Tensor or list of Tensors."""
if isinstance(x, (list, tuple)):
for i in range(len(x)):
x[i][indices] = src[i][indices]
else:
x[indices] = src[indices]
class VOSR(nn.Module):
def __init__(
self,
time_dist = ['lognorm', -0.4, 1.0],
cfg_ratio = 0.10,
cfg_scale = 2.0,
interp_type = 'lin',
u_weight = 1,
a = 1.0,
b = 1.0,
accelerator = None,
t_start: float = 0.0,
t_end: float = 1.0,
args=None,
):
super().__init__()
self.time_dist = time_dist
self.interp_type = interp_type
self.cfg_ratio = cfg_ratio
self.cfg_scale = cfg_scale
self.args = args
self.u_weight = u_weight
self.device = accelerator.device
self.a = a
self.b = b
self.t_start = t_start
self.t_end = t_end
def interpolate(self, cond, uncond, alpha, interp_type='linear'):
if interp_type == 'sph':
return alpha * cond + (1 - alpha**2)**0.5 * uncond
elif interp_type == 'lin':
return alpha * cond + (1 - alpha) * uncond
def sample_t_r_v1(self, B, device):
samples = torch.rand(B, 2, device=device)
t_raw, r_raw = samples[:, 0], samples[:, 1]
t = torch.maximum(t_raw, r_raw)
g = _beta_ratio((B,), self.a, self.b, device)
r = t.detach() * (1. - g)
return t, r
def _prepare_cfg_conditions(self, B, device, lq, z, cond_strength_aelq):
"""Prepare lq_weak/lq_noised/lq_mixed and z_weak/z_noised/z_mixed for CFG training.
Fixed: use_aelq=True, use_venc=True, weak_cond_strength_venc=0, cond_strength_venc=1.
"""
cfg_mask = torch.rand(B, device=device) < self.cfg_ratio
cfg_indices = (cfg_mask > 0).nonzero(as_tuple=True)[0]
weak_cond_strength_aelq = random.uniform(self.args.weak_cond_strength_aelq_list[0], self.args.weak_cond_strength_aelq_list[1])
lq_weak = self.interpolate(lq, torch.zeros_like(lq), weak_cond_strength_aelq, self.interp_type)
lq_noised = self.interpolate(lq, torch.randn_like(lq), cond_strength_aelq, interp_type='sph')
lq_mixed = lq_noised.clone()
lq_mixed[cfg_indices] = lq_weak[cfg_indices]
# weak_cond_strength_venc=0 => z_weak is zeros
z_weak = _zero_like(z)
# cond_strength_venc=1, sph interp with alpha=1 => z_noised = z
z_noised = z
z_mixed = _zero_like(z) # will be overwritten below
if isinstance(z, (list, tuple)):
z_mixed = [zi.clone() for zi in z]
else:
z_mixed = z.clone()
_set_indices(z_mixed, cfg_indices, z_weak)
return cfg_indices, lq_weak, lq_noised, lq_mixed, z_weak, z_noised, z_mixed
def _prepare_cfg_conditions_distill(self, B, device, lq, z, cond_strength_aelq):
"""Same as _prepare_cfg_conditions but uses args.weak_cond_strength_aelq (scalar) for distill losses."""
cfg_mask = torch.rand(B, device=device) < self.cfg_ratio
cfg_indices = (cfg_mask > 0).nonzero(as_tuple=True)[0]
lq_weak = self.interpolate(lq, torch.zeros_like(lq), self.args.weak_cond_strength_aelq, self.interp_type)
lq_noised = self.interpolate(lq, torch.randn_like(lq), cond_strength_aelq, interp_type='sph')
lq_mixed = lq_noised.clone()
lq_mixed[cfg_indices] = lq_weak[cfg_indices]
z_weak = _zero_like(z)
z_noised = z
if isinstance(z, (list, tuple)):
z_mixed = [zi.clone() for zi in z]
else:
z_mixed = z.clone()
_set_indices(z_mixed, cfg_indices, z_weak)
return cfg_indices, lq_weak, lq_noised, lq_mixed, z_weak, z_noised, z_mixed
def _teacher_target(self, model, model_tea, t, t_, lq_noised, lq_weak, z_t, z_noised, z_weak, z_mixed, v_target):
"""Compute teacher/self-distill target v with CFG."""
with torch.no_grad():
inp_cond = torch.cat([lq_noised, z_t], dim=1)
inp_weak = torch.cat([lq_weak, z_t], dim=1)
omega = torch.where((t_ >= self.t_start) & (t_ <= self.t_end), self.cfg_scale, 1.0)
if model_tea is not None:
v_cond = model_tea(inp_cond, t=t, z=z_noised)
v_weak_pred = model_tea(inp_weak, t=t, z=z_weak)
v = v_weak_pred + omega * (v_cond - v_weak_pred)
else:
v_cond = v_target
v_weak_pred = model(inp_weak, t=t, r=t, z=z_weak)
v = v_weak_pred + omega * (v_cond - v_weak_pred)
return v
# ──────────────────── Training losses ──────────────────── #
def loss_fm(self, model, lq, hq, z=None, weight_dtype=None):
B, device = hq.size(0), hq.device
if self.time_dist[0] == 'uniform':
t = torch.rand(B, device=device)
elif self.time_dist[0] == 'lognorm':
mu, sigma = self.time_dist[-2], self.time_dist[-1]
rnd_normal = torch.randn(B, device=device)
t = torch.sigmoid(rnd_normal * sigma + mu)
t_ = rearrange(t, "b -> b 1 1 1")
cond_strength_aelq = torch.sigmoid(torch.randn(B, device=device) * self.args.cond_strength_aelq_list[1] + self.args.cond_strength_aelq_list[0])
cond_strength_aelq = rearrange(cond_strength_aelq, "b -> b 1 1 1")
eps = torch.randn_like(lq)
_, _, _, lq_mixed, _, _, z_mixed = self._prepare_cfg_conditions(B, device, lq, z, cond_strength_aelq)
z_t = (1 - t_) * hq + t_ * eps
inp = torch.cat([lq_mixed, z_t], dim=1)
v = eps - hq
v_current = model(inp, t, z=z_mixed)
v_error = v_current - v
loss = (v_error ** 2).mean()
return loss, loss
def loss_fm_distill_shortcut_improved(self, model, lq, hq, z=None, model_tea=None):
B, device = hq.size(0), hq.device
t, r = self.sample_t_r_v1(B, device)
t_ = rearrange(t, "b -> b 1 1 1")
r_ = rearrange(r, "b -> b 1 1 1")
cond_strength_aelq = torch.sigmoid(torch.randn(B, device=device) * self.args.cond_strength_aelq_list[1] + self.args.cond_strength_aelq_list[0])
cond_strength_aelq = rearrange(cond_strength_aelq, "b -> b 1 1 1")
eps = torch.randn_like(lq)
_, lq_weak, lq_noised, lq_mixed, z_weak, z_noised, z_mixed = self._prepare_cfg_conditions_distill(B, device, lq, z, cond_strength_aelq)
z_t = (1 - t_) * hq + t_ * eps
v_target = eps - hq
inp_mixed = torch.cat([lq_mixed, z_t], dim=1)
v = self._teacher_target(model, model_tea, t, t_, lq_noised, lq_weak, z_t, z_noised, z_weak, z_mixed, v_target)
v_current = model(inp_mixed, t=t, r=t, z=z_mixed)
v_loss = ((v_current - v) ** 2).mean()
with torch.no_grad():
s = (t + r) / 2
inp_i = torch.cat([lq_mixed, z_t], dim=1)
v1 = model(inp_i, t=t, r=s, z=z_mixed)
z_s = z_t - (t_-r_)/2*v1
inp_i = torch.cat([lq_mixed, z_s], dim=1)
v2 = model(inp_i, t=s, r=r, z=z_mixed)
integral_approx = (v1 + v2) / 2
u_current = model(inp_mixed, t, r, z_mixed)
u_loss = ((u_current - stopgrad(integral_approx)) ** 2).mean()
loss = v_loss + self.u_weight * u_loss
return loss, loss
def _rcgm_consistency(self, model, rng_state, inp_mixed, lq_mixed, z_t, z_mixed, t, r, t_, r_, v, DELTA_T, RCGM_N):
"""Compute RCGM consistency loss."""
torch.cuda.set_rng_state(rng_state)
u_current = model(inp_mixed, t=t, r=r, z=z_mixed)
with torch.no_grad():
delta = DELTA_T
t_m_ = (t_ - delta).clamp(min=r_)
actual_delta = t_ - t_m_
z_t_m = z_t - v * actual_delta
z_cur = z_t_m
Ft_tar = torch.zeros_like(z_t)
for step_i in range(RCGM_N):
frac_start = step_i / RCGM_N
frac_end = (step_i + 1) / RCGM_N
t_step_ = t_m_ * (1 - frac_start) + r_ * frac_start
t_next_ = t_m_ * (1 - frac_end) + r_ * frac_end
dt = t_step_ - t_next_
torch.cuda.set_rng_state(rng_state)
inp_step = torch.cat([lq_mixed, z_cur], dim=1)
v_step = model(inp_step, t=t_step_.reshape(-1), r=t_next_.reshape(-1), z=z_mixed)
z_cur = z_cur - dt * v_step
Ft_tar = Ft_tar + v_step * dt
d = (t_ - r_).clamp(min=1e-6)
near_mask = (t_ < (r_ + delta + 1e-6))
cof_l = torch.where(near_mask, torch.ones_like(t_), d / delta)
cof_r = torch.where(near_mask, 1.0 / d, torch.ones_like(t_) / delta)
correction = (u_current.detach() * cof_l - Ft_tar * cof_r) - v
correction = correction.clamp(min=-1.0, max=1.0)
rcgm_target = u_current.detach() - correction
u_loss = ((u_current - stopgrad(rcgm_target)) ** 2).mean()
return u_loss
def loss_fm_distill_rcgm_improved(self, model, lq, hq, z=None, model_tea=None):
B, device = hq.size(0), hq.device
DELTA_T = getattr(self.args, 'rcgm_delta_t', 0.01)
RCGM_N = getattr(self.args, 'rcgm_n_steps', 2)
t, r = self.sample_t_r_v1(B, device)
t_ = rearrange(t, "b -> b 1 1 1")
r_ = rearrange(r, "b -> b 1 1 1")
cond_strength_aelq = torch.sigmoid(torch.randn(B, device=device) * self.args.cond_strength_aelq_list[1] + self.args.cond_strength_aelq_list[0])
cond_strength_aelq = rearrange(cond_strength_aelq, "b -> b 1 1 1")
eps = torch.randn_like(lq)
_, lq_weak, lq_noised, lq_mixed, z_weak, z_noised, z_mixed = self._prepare_cfg_conditions_distill(B, device, lq, z, cond_strength_aelq)
z_t = (1 - t_) * hq + t_ * eps
v_target = eps - hq
inp_mixed = torch.cat([lq_mixed, z_t], dim=1)
v = self._teacher_target(model, model_tea, t, t_, lq_noised, lq_weak, z_t, z_noised, z_weak, z_mixed, v_target)
rng_state = torch.cuda.get_rng_state()
v_current = model(inp_mixed, t=t, r=t, z=z_mixed)
v_loss = ((v_current - v) ** 2).mean()
u_loss = self._rcgm_consistency(model, rng_state, inp_mixed, lq_mixed, z_t, z_mixed, t, r, t_, r_, v, DELTA_T, RCGM_N)
loss = v_loss + self.u_weight * u_loss
return loss, loss
# ──────────────────── Sampling ──────────────────── #
@torch.no_grad()
def sample_onestep(self, model, lq, venc_fea=None, n_steps: int = 8, schedule: str = "linear"):
B, device = lq.size(0), lq.device
z = torch.randn_like(lq)
if schedule == "linear":
t_seq = torch.linspace(1., 0., n_steps + 1, device=device)
elif schedule == "cosine":
t_seq = torch.cos(torch.linspace(0, np.pi / 2, n_steps + 1, device=device))
else:
raise ValueError("schedule must be 'linear' or 'cosine'")
for i in range(n_steps):
t_cur, t_next = t_seq[i], t_seq[i + 1]
inp = torch.cat([lq, z], dim=1)
u = model(inp, t_cur.repeat(B), t_next.repeat(B), venc_fea)
z = z - (t_cur - t_next) * u
return z
@torch.no_grad()
def sample_multistep_fm(self, model, lq, venc_fea=None, n_steps: int = 25, schedule: str = "linear"):
B, device = lq.size(0), lq.device
eps = torch.randn_like(lq)
# CFG condition preparation (cfg_ratio > 0 always)
weak_cond_strength_aelq = (self.args.weak_cond_strength_aelq_list[0] + self.args.weak_cond_strength_aelq_list[1]) / 2.
lq_weak = self.interpolate(lq, torch.zeros_like(lq), weak_cond_strength_aelq, self.interp_type)
# weak_cond_strength_venc=0 => venc_fea_weak is zeros; cond_strength_venc=1 => venc_fea unchanged
if isinstance(venc_fea, (list, tuple)):
venc_fea_weak = [torch.zeros_like(fea) for fea in venc_fea]
else:
venc_fea_weak = torch.zeros_like(venc_fea) if venc_fea is not None else None
z = eps
if schedule == "linear":
t_seq = torch.linspace(1., 0., n_steps + 1, device=device)
elif schedule == "cosine":
t_seq = torch.cos(torch.linspace(0, np.pi / 2, n_steps + 1, device=device))
else:
raise ValueError("schedule must be 'linear' or 'cosine'")
for i in range(n_steps):
t_cur, t_next = t_seq[i], t_seq[i + 1]
dt = t_cur - t_next
model_inp = torch.cat([lq, z], dim=1)
model_t = t_cur.repeat(B)
model_venc_fea = venc_fea
if t_cur <= self.t_end and t_cur >= self.t_start:
inp_weak = torch.cat([lq_weak, z], dim=1)
model_inp = torch.cat([model_inp, inp_weak], dim=0)
model_t = model_t.repeat(2)
if venc_fea is not None:
if isinstance(venc_fea, (list, tuple)):
model_venc_fea = [torch.cat([v, w], dim=0) for v, w in zip(venc_fea, venc_fea_weak)]
else:
model_venc_fea = torch.cat([venc_fea, venc_fea_weak], dim=0)
d_cur = model(model_inp, model_t, z=model_venc_fea)
if t_cur <= self.t_end and t_cur >= self.t_start:
d_cur_cond, d_cur_weak = d_cur.chunk(2)
d_cur = d_cur_weak + self.cfg_scale * (d_cur_cond - d_cur_weak)
z = z - dt * d_cur
return z