upgrade from 1.2.0 to 1.3.0

This commit is contained in:
maochaojie
2024-11-19 19:20:02 +08:00
parent 4ec0492897
commit a683061c6f
87 changed files with 10400 additions and 1117 deletions
+17 -39
View File
@@ -19,10 +19,6 @@ class BaseDiffusion(object):
para_dict = {
'NOISE_SCHEDULER': {},
'SAMPLER_SCHEDULER': {},
'MIN_SNR_GAMMA': {
'value': None,
'description': 'The minimum SNR gamma value for the loss function.'
},
'PREDICTION_TYPE': {
'value': 'eps',
'description':
@@ -37,7 +33,6 @@ class BaseDiffusion(object):
self.init_params()
def init_params(self):
self.min_snr_gamma = self.cfg.get('MIN_SNR_GAMMA', None)
self.prediction_type = self.cfg.get('PREDICTION_TYPE', 'eps')
self.noise_scheduler = NOISE_SCHEDULERS.build(self.cfg.NOISE_SCHEDULER,
logger=self.logger)
@@ -67,17 +62,19 @@ class BaseDiffusion(object):
show_progress=False,
return_intermediate=None,
intermediate_callback=None,
reverse_scale = -1.,
x = None,
**kwargs):
assert isinstance(steps, (int, torch.LongTensor))
assert return_intermediate in (None, 'x0', 'xt')
assert isinstance(sampler, (str, dict, Config))
intermediates = []
def callback_fn(x_t, t, sigma=None, alpha=None):
def callback_fn(x_t, t, sigma=None, alpha_bar=None):
timestamp = t
t = t.repeat(len(x_t)).round().long().to(x_t.device)
sigma = sigma.repeat(len(x_t), *([1] * (len(sigma.shape) - 1)))
alpha = alpha.repeat(len(x_t), *([1] * (len(alpha.shape) - 1)))
alpha_bar = alpha_bar.repeat(len(x_t), *([1] * (len(alpha_bar.shape) - 1)))
if guide_scale is None or guide_scale == 1.0:
out = model(x=x_t, t=t, **model_kwargs)
@@ -101,15 +98,12 @@ class BaseDiffusion(object):
if self.prediction_type == 'x0':
x0 = out
elif self.prediction_type == 'eps':
x0 = (x_t - sigma * out) / alpha
x0 = (x_t - sigma * out) / alpha_bar
elif self.prediction_type == 'v':
x0 = alpha * x_t - sigma * out
x0 = alpha_bar * x_t - sigma * out
else:
raise NotImplementedError(
f'prediction_type {self.prediction_type} not implemented')
# print("torch.sum(y_out):", torch.sum(y_out), "torch.sum(u_out):", torch.sum(u_out), "torch.sum(out):",
# torch.sum(out), "torch.sum(x0):", torch.sum(x0), "sigmas", sigma, "alphas", alpha)
return x0
sampler_ins = self.get_sampler(sampler)
@@ -117,12 +111,14 @@ class BaseDiffusion(object):
# this is ignored for schnell
sampler_output = sampler_ins.preprare_sampler(
noise,
x = x,
steps=steps,
reverse_scale= reverse_scale,
prediction_type=self.prediction_type,
scheduler_ins=self.sampler_scheduler,
callback_fn=callback_fn)
for _ in trange(steps, disable=not show_progress):
for _ in trange(sampler_output.steps, disable=not show_progress):
trange.desc = sampler_output.msg
sampler_output = sampler_ins.step(sampler_output)
if return_intermediate == 'x_0':
@@ -145,30 +141,19 @@ class BaseDiffusion(object):
if noise is None:
noise = torch.randn_like(x_0)
schedule_output = self.noise_scheduler.add_noise(x_0, noise, **kwargs)
x_t, t, sigma, alpha = schedule_output.x_t, schedule_output.t, schedule_output.sigma, schedule_output.alpha
x_t, t, sigma, alpha_bar = schedule_output.x_t, schedule_output.t, schedule_output.sigma, schedule_output.alpha_bar
out = model(x=x_t, t=t, **model_kwargs)
# mse loss
target = {
'eps': noise,
'x0': x_0,
'v': alpha * noise - sigma * x_0
'v': alpha_bar * noise - sigma * x_0
}[self.prediction_type]
loss = (out - target).pow(2)
if reduction == 'mean':
loss = loss.flatten(1).mean(dim=1)
if self.min_snr_gamma is not None:
alphas = self.noise_scheduler.alphas.to(x_0.device)[t]
sigmas = self.noise_scheduler.sigmas.pow(2).to(x_0.device)[t]
snrs = (alphas / sigmas).clamp(min=1e-20)
min_snrs = snrs.clamp(max=self.min_snr_gamma)
weights = min_snrs / snrs
else:
weights = 1
loss = loss * weights
return loss
def get_sampler(self, sampler):
@@ -248,17 +233,6 @@ class DiffusionFluxRF(BaseDiffusion):
loss = (target - out)**2
if reduction == 'mean':
loss = loss.flatten(1).mean(dim=1)
if self.min_snr_gamma is not None:
alphas = self.noise_scheduler.alphas.to(x_0.device)[t]
sigmas = self.noise_scheduler.sigmas.pow(2).to(x_0.device)[t]
snrs = (alphas / sigmas).clamp(min=1e-20)
min_snrs = snrs.clamp(max=self.min_snr_gamma)
weights = min_snrs / snrs
else:
weights = 1
loss = loss * weights
return loss
@torch.no_grad()
@@ -271,6 +245,8 @@ class DiffusionFluxRF(BaseDiffusion):
show_progress=False,
return_intermediate=None,
intermediate_callback=None,
reverse_scale=-1.,
x=None,
**kwargs):
# sanity check
assert isinstance(steps, (int, torch.LongTensor))
@@ -278,7 +254,7 @@ class DiffusionFluxRF(BaseDiffusion):
assert isinstance(sampler, (str, dict, Config))
intermediates = []
def callback_fn(x_t, t, sigma=None, alpha=None):
def callback_fn(x_t, t, sigma=None, alpha_bar=None):
sigma = torch.full((x_t.shape[0], ),
sigma,
dtype=x_t.dtype,
@@ -291,12 +267,14 @@ class DiffusionFluxRF(BaseDiffusion):
# this is ignored for schnell
sampler_output = sampler_ins.preprare_sampler(
noise,
x=x,
steps=steps,
reverse_scale=reverse_scale,
prediction_type=self.prediction_type,
scheduler_ins=self.sampler_scheduler,
callback_fn=callback_fn)
for _ in trange(steps, disable=not show_progress):
for _ in trange(sampler_output.steps, disable=not show_progress):
trange.desc = sampler_output.msg
sampler_output = sampler_ins.step(sampler_output)
if return_intermediate == 'x_0':
+71 -17
View File
@@ -15,15 +15,18 @@ class SamplerOutput(object):
callback_fn: callable
prediction_type: str
alphas: torch.Tensor
alphas_bar: torch.Tensor
betas: torch.Tensor
sigmas: torch.Tensor
alphas_init: torch.Tensor
alphas_bar_init: torch.Tensor
betas_init: torch.Tensor
sigmas_init: torch.Tensor
ts: torch.Tensor
x_t: torch.Tensor
x_0: torch.Tensor
step: int
steps: int
msg: str
def add_custom_field(self, key: str, value) -> None:
@@ -49,7 +52,7 @@ class BaseDiffusionSampler(object):
self.t_max = self.cfg.get('T_MAX', None)
self.t_min = self.cfg.get('T_MIN', None)
def discretization(self, steps=20, num_timesteps=1000, **kwargs):
def discretization(self, steps=20, num_timesteps=1000, reverse_scale = -1., **kwargs):
# get timesteps
if isinstance(steps, int):
steps += 1 if self.discard_penultimate_step else 0
@@ -74,17 +77,23 @@ class BaseDiffusionSampler(object):
steps = steps.clamp_(t_min, t_max)
elif isinstance(steps, list):
steps = torch.tensor(steps)
timesteps = torch.as_tensor(steps, dtype=torch.float32)
return timesteps
if reverse_scale >=0:
img2img_step = int((1 - reverse_scale) * len(steps))
timesteps = torch.as_tensor(steps[img2img_step:], dtype=torch.float32)
return timesteps
return torch.as_tensor(steps, dtype=torch.float32)
def preprare_sampler(self,
noise,
x=None,
steps=20,
reverse_scale=-1.,
scheduler_ins=None,
prediction_type='',
sigmas=None,
betas=None,
alphas=None,
alphas_bar=None,
callback_fn=None,
**kwargs):
'''
@@ -96,36 +105,52 @@ class BaseDiffusionSampler(object):
4. To ensure the safety of threading, use the instance of SamplerOutput as the manager,
which manage all necessary information.
'''
if reverse_scale >= 0:
assert x is not None
num_timesteps = scheduler_ins.num_timesteps if scheduler_ins is not None else 1000
timestamps = self.discretization(steps,
num_timesteps=num_timesteps,
reverse_scale=reverse_scale,
**kwargs)
alphas = scheduler_ins.t_to_alpha(
timestamps, **kwargs) if scheduler_ins is not None else alphas
alphas_bar = scheduler_ins.t_to_alpha_bar(
timestamps, **kwargs) if scheduler_ins is not None else alphas_bar
betas = scheduler_ins.t_to_beta(
timestamps, **kwargs) if scheduler_ins is not None else betas
sigmas = scheduler_ins.t_to_sigma(
timestamps, **kwargs) if scheduler_ins is not None else sigmas
alphas_init = scheduler_ins.t_to_alpha_init(
timestamps, **kwargs) if scheduler_ins is not None else alphas
alphas_bar_init = scheduler_ins.t_to_alpha_bar_init(
timestamps, **kwargs) if scheduler_ins is not None else alphas_bar
betas_init = scheduler_ins.t_to_beta_init(
timestamps, **kwargs) if scheduler_ins is not None else betas
sigmas_init = scheduler_ins.t_to_sigma_init(
timestamps, **kwargs) if scheduler_ins is not None else sigmas
if reverse_scale >= 0:
x_t = x_0 = scheduler_ins.add_noise(x, noise=noise, t=timestamps[0].repeat(x.size(0)).to(x.device)).x_t if len(timestamps) > 0 else x
else:
x_t = x_0 = noise
# Consider the sigma's list is from sigma_ to zero. the steps equal to len(timestamps)
output = SamplerOutput(callback_fn=callback_fn,
prediction_type=prediction_type,
alphas=alphas,
alphas_bar=alphas_bar,
betas=betas,
sigmas=sigmas,
alphas_init=alphas_init,
alphas_bar_init=alphas_bar_init,
betas_init=betas_init,
sigmas_init=sigmas_init,
ts=timestamps,
x_t=noise,
x_0=noise,
x_t=x_t,
x_0=x_0,
step=0,
msg='step 0')
msg='step 0',
steps=len(timestamps) - 1)
return output
def step(self, sampler_ouput):
@@ -159,22 +184,35 @@ class DDIMSampler(BaseDiffusionSampler):
def preprare_sampler(self,
noise,
x=None,
steps=20,
reverse_scale = -1.,
scheduler_ins=None,
prediction_type='',
sigmas=None,
betas=None,
alphas=None,
alphas_bar=None,
callback_fn=None,
**kwargs):
output = super().preprare_sampler(noise, steps, scheduler_ins,
prediction_type, sigmas, betas,
alphas, callback_fn, **kwargs)
output = super().preprare_sampler(noise,
x = x,
steps = steps,
reverse_scale = reverse_scale,
scheduler_ins = scheduler_ins,
prediction_type = prediction_type,
sigmas = sigmas,
betas = betas,
alphas = alphas,
alphas_bar = alphas_bar,
callback_fn = callback_fn,
**kwargs)
sigmas = output.sigmas
sigmas = torch.cat([sigmas, sigmas.new_zeros([1])])
sigmas_vp = (sigmas**2 / (1 + sigmas**2))**0.5
sigmas_vp[sigmas == float('inf')] = 1.
output.add_custom_field('sigmas_vp', sigmas_vp)
output.steps += 1
return output
def step(self, sampler_output):
@@ -182,10 +220,10 @@ class DDIMSampler(BaseDiffusionSampler):
step = sampler_output.step
t = sampler_output.ts[step]
sigmas_vp = sampler_output.sigmas_vp.to(x_t.device)
alpha_init = _i(sampler_output.alphas_init, step, x_t[:1])
alpha_bar_init = _i(sampler_output.alphas_bar_init, step, x_t[:1])
sigma_init = _i(sampler_output.sigmas_init, step, x_t[:1])
x = sampler_output.callback_fn(x_t, t, sigma_init, alpha_init)
x = sampler_output.callback_fn(x_t, t, sigma_init, alpha_bar_init)
noise_factor = self.eta * (sigmas_vp[step + 1]**2 /
sigmas_vp[step]**2 *
(1 - (1 - sigmas_vp[step]**2) /
@@ -202,16 +240,19 @@ class DDIMSampler(BaseDiffusionSampler):
return sampler_output
@DIFFUSION_SAMPLERS.register_class('flow_eluer')
@DIFFUSION_SAMPLERS.register_class('flow_euler')
class FlowEluerSampler(BaseDiffusionSampler):
def preprare_sampler(self,
noise,
x=None,
steps=20,
reverse_scale = -1.,
scheduler_ins=None,
prediction_type='',
sigmas=None,
betas=None,
alphas=None,
alphas_bar=None,
callback_fn=None,
**kwargs):
if noise.ndim == 3:
@@ -220,9 +261,18 @@ class FlowEluerSampler(BaseDiffusionSampler):
n, _, h, w = noise.shape
seq_len = (h // 2 * w // 2)
kwargs['seq_len'] = seq_len
output = super().preprare_sampler(noise, steps, scheduler_ins,
prediction_type, sigmas, betas,
alphas, callback_fn, **kwargs)
output = super().preprare_sampler(noise,
x = x,
steps = steps,
reverse_scale = reverse_scale,
scheduler_ins = scheduler_ins,
prediction_type = prediction_type,
sigmas = sigmas,
betas = betas,
alphas = alphas,
alphas_bar = alphas_bar,
callback_fn = callback_fn,
**kwargs)
return output
def step(self, sampler_output):
@@ -241,9 +291,13 @@ class FlowEluerSampler(BaseDiffusionSampler):
sampler_output.msg = f'step {step}, sigma_curr: {sigma_curr}, sigma_prev: {sigma_prev}'
return sampler_output
def discretization(self, steps=20, num_timesteps=1000, **kwargs):
def discretization(self, steps=20, num_timesteps=1000, reverse_scale=-1., **kwargs):
# extra step for zero
timesteps = torch.linspace(num_timesteps, 0, steps + 1)
if reverse_scale >= 0:
img2img_step = int((1 - reverse_scale) * len(timesteps))
timesteps = timesteps[img2img_step:]
return timesteps
return timesteps
@staticmethod
+52 -11
View File
@@ -21,7 +21,7 @@ class ScheduleOutput(object):
x_0: torch.Tensor
t: torch.Tensor
sigma: torch.Tensor
alpha: torch.Tensor
alpha_bar: torch.Tensor
custom_fields: dict = field(default_factory=dict)
def add_custom_field(self, key: str, value) -> None:
@@ -30,6 +30,21 @@ class ScheduleOutput(object):
@NOISE_SCHEDULERS.register_class()
class BaseNoiseScheduler(object):
'''
In the diffusion model, the parameters related to the noise schedule are alpha, beta,
and sigma. The following are the definitions of the above three parameters, which should
be the basic property for the instance of noise scheduler.
\alpha_{t} = \sqrt{1 - \beta_{t}^2} \alpha is the strength of signal and \beta is the strength of noise
\sigma_{t} = \sqrt{1 - \overline\alpha} = \sqrt{1 - \prod_{i=1}^{t}\alpha^2_{i}} (P(x_{t}|x_{0}) ~ N(\overline\alpha x_{0}, \sigma^2))
\alpha_bar_{t} = \sqrt{\overline\alpha} = \sqrt{\prod_{i=1}^{t}\alpha^2_{i}} (P(x_{t}|x_{0}) ~ N(\overline\alpha x_{0}, \sigma^2))
where sigma_{t} is the var of p(x_{t-1}|x_{t}, x_{0}).
(reference to https://arxiv.org/abs/2010.02502)
let sigma transfer to beta:
square_\beta = 1 - \frac{1 - square_\sigma_{t}}{1 - square_\sigma_{t - 1 }}
'''
para_dict = {
'NUM_TIMESTEPS': {
'value': 1000,
@@ -48,7 +63,7 @@ class BaseNoiseScheduler(object):
self.num_timesteps = self.cfg.get('NUM_TIMESTEPS', 1000)
self._sample_steps = torch.arange(self.num_timesteps,
dtype=torch.float32)
self._sigmas, self._betas, self._alphas, self._timesteps = None, None, None, None
self._sigmas, self._betas, self._alphas, self._alphas_bar, self._timesteps = None, None, None, None, None
def check_function(self):
try:
@@ -128,6 +143,10 @@ class BaseNoiseScheduler(object):
square_beta = self.sigmas_to_square_betas(sigma)
return torch.sqrt(1 - square_beta)
def t_to_alpha_bar(self, t, **kwargs):
sigma = self.t_to_sigma(t)
return torch.sqrt(1 - sigma**2)
def t_to_beta(self, t, **kwargs):
sigma = self.t_to_sigma(t)
square_beta = self.sigmas_to_square_betas(sigma)
@@ -138,11 +157,11 @@ class BaseNoiseScheduler(object):
t = torch.randint(0,
self.num_timesteps, (x_0.shape[0], ),
device=x_0.device).long()
alpha = _i(self.alphas, t, x_0)
alpha = _i(self.alphas_bar, t, x_0)
sigma = _i(self.sigmas, t, x_0)
x_t = alpha * x_0 + sigma * noise
return ScheduleOutput(x_0=x_0, x_t=x_t, t=t, alpha=alpha, sigma=sigma)
return ScheduleOutput(x_0=x_0, x_t=x_t, t=t, alpha_bar=alpha, sigma=sigma)
def t_to_alpha_init(self, t, **kwargs):
indices = t.long()
@@ -153,6 +172,16 @@ class BaseNoiseScheduler(object):
alpha = self.alphas[step_indices].flatten().to(t)
return alpha
def t_to_alpha_bar_init(self, t, **kwargs):
indices = t.long()
indices[indices >= self.num_timesteps] = self.num_timesteps - 1
timesteps = self.timesteps.to(t)[indices]
step_indices = [(self.timesteps.to(t) == t).nonzero().item()
for t in timesteps]
alpha_bar = self.alphas_bar[step_indices].flatten().to(t)
return alpha_bar
def t_to_beta_init(self, t, **kwargs):
indices = t.long()
indices[indices >= self.num_timesteps] = self.num_timesteps - 1
@@ -205,6 +234,10 @@ class BaseNoiseScheduler(object):
def alphas(self):
return self._alphas
@property
def alphas_bar(self):
return self._alphas_bar
@property
def timesteps(self):
return self._timesteps
@@ -221,6 +254,10 @@ class BaseNoiseScheduler(object):
'data': self._alphas.cpu().numpy(),
'label': 'alphas'
}, {
'data': self._alphas_bar.cpu().numpy(),
'label': 'alphas_bar'
},
{
'data': self._timesteps.cpu().numpy() / self.num_timesteps,
'label': 'timesteps'
}]
@@ -280,7 +317,8 @@ class ScaledLinearScheduler(BaseNoiseScheduler):
self.snr_shift_scale,
self.rescale_betas_zero_snr)
self._betas = torch.sqrt(square_betas)
self._alphas = torch.sqrt(1 - self._sigmas**2)
self._alphas = torch.sqrt(1 - square_betas)
self._alphas_bar = torch.sqrt(1 - self._sigmas**2)
self._timesteps = torch.arange(len(self._sigmas), dtype=torch.float32)
@@ -304,7 +342,8 @@ class LinearScheduler(BaseNoiseScheduler):
sigmas = self.betas_to_sigmas(betas)
self._sigmas = sigmas
self._betas = betas
self._alphas = torch.sqrt(1 - sigmas**2)
self._alphas = torch.sqrt(1 - betas**2)
self._alphas_bar = torch.sqrt(1 - sigmas**2)
self._timesteps = torch.arange(len(sigmas), dtype=torch.float32)
@@ -319,7 +358,8 @@ class FlowMatchUniformScheduler(BaseNoiseScheduler):
self._timesteps = timesteps
self._sigmas = self.t_to_sigma(timesteps)
self._betas = torch.sqrt(self.sigmas_to_square_betas(self._sigmas))
self._alphas = torch.sqrt(1 - self.betas**2)
self._alphas = torch.sqrt(1 - self._betas**2)
self._alphas_bar = torch.sqrt(1 - self._sigmas ** 2)
def add_noise(self, x_0, noise=None, t=None, **kwargs):
if t is None:
@@ -332,7 +372,7 @@ class FlowMatchUniformScheduler(BaseNoiseScheduler):
x_t=x_t,
t=t,
sigma=sigma,
alpha=self.t_to_alpha(t))
alpha_bar=self.t_to_alpha_bar(t))
def sigma_to_t(self, sigma, **kwargs):
return sigma * self.num_timesteps
@@ -406,7 +446,7 @@ class FlowMatchShiftScheduler(FlowMatchUniformScheduler):
x_t=x_t,
t=t,
sigma=sigma,
alpha=self.t_to_alpha(t))
alpha_bar=self.t_to_alpha_bar(t))
def sigma_to_t(self, sigma, **kwargs):
t = sigma / (sigma - self.shift * sigma + self.shift)
@@ -486,7 +526,7 @@ class FlowMatchFluxShiftScheduler(FlowMatchUniformScheduler):
x_t=x_t,
t=t,
sigma=sigma,
alpha=self.t_to_alpha(t))
alpha_bar=self.t_to_alpha_bar(t))
def sigma_to_t(self, sigma, **kwargs):
seq_len = kwargs.get('seq_len', 256)
@@ -570,6 +610,7 @@ class FlowMatchSigmaScheduler(FlowMatchUniformScheduler):
(self.shift - 1) * timesteps)
self._betas = torch.sqrt(self.sigmas_to_square_betas(self._sigmas))
self._alphas = torch.sqrt(1 - self.betas**2)
self._alphas_bar = torch.sqrt(1 - self._sigmas ** 2)
def add_noise(self, x_0, noise=None, t=None, **kwargs):
if t is None:
@@ -589,7 +630,7 @@ class FlowMatchSigmaScheduler(FlowMatchUniformScheduler):
x_t=x_t,
t=t,
sigma=sigma,
alpha=self.t_to_alpha(t))
alpha_bar=self.t_to_alpha_bar(t))
def compute_density_for_timestep_sampling(self, t):
"""Compute the density for sampling the timesteps when doing SD3 training.