upgrade from 1.2.0 to 1.3.0
This commit is contained in:
@@ -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':
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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.
|
||||
|
||||
Reference in New Issue
Block a user