Files
modelscope-scepter/scepter/modules/model/diffusion/schedules.py
T
2025-02-03 13:36:44 +08:00

714 lines
27 KiB
Python

# -*- coding: utf-8 -*-
# Copyright (c) Alibaba, Inc. and its affiliates.
import math
import random
from dataclasses import dataclass, field
from typing import Callable
import numpy as np
import torch
from torch import Tensor
from scepter.modules.model.registry import NOISE_SCHEDULERS
from scepter.modules.utils.config import dict_to_yaml
from scepter.modules.utils.math_plot import plot_multi_curves
from .util import _i
@dataclass
class ScheduleOutput(object):
x_t: torch.Tensor
x_0: torch.Tensor
t: torch.Tensor
sigma: torch.Tensor
alpha_bar: torch.Tensor
custom_fields: dict = field(default_factory=dict)
def add_custom_field(self, key: str, value) -> None:
self.__setattr__(key, value)
@NOISE_SCHEDULERS.register_class()
class BaseNoiseScheduler(object):
r'''
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,
'description': 'The number of timesteps for sampling.'
},
}
def __init__(self, cfg, logger=None):
super(BaseNoiseScheduler, self).__init__()
self.logger = logger
self.cfg = cfg
self.init_params()
self.get_schedule()
def init_params(self):
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._alphas_bar, self._timesteps = None, None, None, None, None
def check_function(self):
try:
predict_timestamps = self.sigma_to_t(self.sigmas)
predict_sigmas = self.t_to_sigma(self._timesteps)
diff_sigmas = torch.sum(torch.abs(predict_sigmas - self.sigmas))
diff_timestamps = torch.sum(
torch.abs(predict_timestamps - self._timesteps))
if diff_sigmas > 1e-3 or diff_timestamps > 1:
self.logger.info(
f'The noise scheduler {self.__class__.__name__} is not correct, '
f'please check the function sigma_to_t or t_to_sigma.'
f'Info: diff sigmas {diff_sigmas}, diff timestamps {diff_timestamps}'
)
raise 'The noise scheduler checked failed.'
else:
self.logger.info(
f'The noise scheduler {self.__class__.__name__} is checked and passed.'
)
except Exception as e:
if isinstance(e, NotImplementedError):
self.logger.info(
'Not implemented function sigma_to_t or t_to_sigma, skip check.'
)
else:
self.logger.info(
f'The noise scheduler {self.__class__.__name__} is not correct, '
f'please check the function sigma_to_t or t_to_sigma. Error: {e}'
)
raise e
def get_schedule(self):
raise NotImplementedError(
'NoiseScheduler get_schedule function not implemented')
def square_betas_to_sigmas(self, square_betas):
return torch.sqrt(1 - torch.cumprod(1 - square_betas, dim=0))
def sigmas_to_square_betas(self, sigmas):
square_alphas = 1 - sigmas**2
betas = 1 - torch.cat(
[square_alphas[:1], square_alphas[1:] / square_alphas[:-1]])
return betas
def sigma_to_t(self, sigma, **kwargs):
if sigma == float('inf'):
t = torch.full_like(sigma, len(self._sigmas) - 1)
else:
log_sigmas = torch.sqrt(self._sigmas**2 /
(1 - self._sigmas**2)).log().to(sigma)
log_sigma = sigma.log()
dists = log_sigma - log_sigmas[:, None]
low_idx = dists.ge(0).cumsum(dim=0).argmax(dim=0).clamp(
max=log_sigmas.shape[0] - 2)
high_idx = low_idx + 1
low, high = log_sigmas[low_idx], log_sigmas[high_idx]
w = (low - log_sigma) / (low - high)
w = w.clamp(0, 1)
t = (1 - w) * low_idx + w * high_idx
t = t.view(sigma.shape)
if t.ndim == 0:
t = t.unsqueeze(0)
return t
def t_to_sigma(self, t, **kwargs):
t = t.float()
low_idx, high_idx, w = t.floor().long(), t.ceil().long(), t.frac()
log_sigmas = torch.sqrt(self.sigmas**2 /
(1 - self.sigmas**2)).log().to(t)
log_sigma = (1 - w) * log_sigmas[low_idx] + w * log_sigmas[high_idx]
log_sigma[torch.isnan(log_sigma)
| torch.isinf(log_sigma)] = float('inf')
return log_sigma.exp()
def t_to_alpha(self, t, **kwargs):
sigma = self.t_to_sigma(t)
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)
return torch.sqrt(square_beta)
def add_noise(self, x_0, noise=None, t=None, **kwargs):
if t is None:
t = torch.randint(0,
self.num_timesteps, (x_0.shape[0], ),
device=x_0.device).long()
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_bar=alpha, sigma=sigma)
def t_to_alpha_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 = 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
timesteps = self.timesteps.to(t)[indices]
step_indices = [(self.timesteps.to(t) == t).nonzero().item()
for t in timesteps]
beta = self.betas[step_indices].flatten().to(t)
return beta
def t_to_sigma_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]
sigma = self.sigmas[step_indices].flatten().to(t)
return sigma
def rescale_zero_terminal_snr(self, alphas_cumprod):
"""
Rescales betas to have zero terminal SNR Based on https://arxiv.org/pdf/2305.08891.pdf (Algorithm 1)
Args:
betas (`torch.Tensor`):
the betas that the scheduler is being initialized with.
Returns:
`torch.Tensor`: rescaled betas with zero terminal SNR
"""
alphas_bar_sqrt = alphas_cumprod.sqrt()
# Store old values.
alphas_bar_sqrt_0 = alphas_bar_sqrt[0].clone()
alphas_bar_sqrt_T = alphas_bar_sqrt[-1].clone()
# Shift so the last timestep is zero.
alphas_bar_sqrt -= alphas_bar_sqrt_T
# Scale so the first timestep is back to the old value.
alphas_bar_sqrt *= alphas_bar_sqrt_0 / (alphas_bar_sqrt_0 -
alphas_bar_sqrt_T)
# Convert alphas_bar_sqrt to betas
alphas_bar = alphas_bar_sqrt**2 # Revert sqrt
return alphas_bar
@property
def sigmas(self):
return self._sigmas
@property
def betas(self):
return self._betas
@property
def alphas(self):
return self._alphas
@property
def alphas_bar(self):
return self._alphas_bar
@property
def timesteps(self):
return self._timesteps
# plot the noise sampling map
def plot_noise_sampling_map(self, save_path):
y = [{
'data': self._sigmas.cpu().numpy(),
'label': 'sigmas'
}, {
'data': self._betas.cpu().numpy(),
'label': 'betas'
}, {
'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'
}]
plot_multi_curves(
x=self._sample_steps.cpu().numpy(),
y=y,
x_label='timesteps',
y_label=None,
title=f"{self.__class__.__name__}'s noise sampling map",
save_path=save_path)
return save_path
def __repr__(self) -> str:
return f'{self.__class__.__name__}' + ' ' + super().__repr__()
@staticmethod
def get_config_template():
return dict_to_yaml('NOISE_SCHEDULER',
__class__.__name__,
BaseNoiseScheduler.para_dict,
set_name=True)
@NOISE_SCHEDULERS.register_class()
class ScaledLinearScheduler(BaseNoiseScheduler):
para_dict = {}
def init_params(self):
super().init_params()
self.beta_min = self.cfg.get('BETA_MIN', 0.00085)
self.beta_max = self.cfg.get('BETA_MAX', 0.012)
self.snr_shift_scale = self.cfg.get('SNR_SHIFT_SCALE', None)
self.rescale_betas_zero_snr = self.cfg.get('RESCALE_BETAS_ZERO_SNR',
False)
def square_betas_to_sigmas(self,
square_betas,
snr_shift_scale=None,
rescale_betas_zero_snr=False):
if snr_shift_scale is not None or rescale_betas_zero_snr:
alphas_cumprod = torch.cumprod(1 - square_betas, dim=0)
if snr_shift_scale is not None and snr_shift_scale > 0:
alphas_cumprod = alphas_cumprod / (
snr_shift_scale + (1 - snr_shift_scale) * alphas_cumprod)
if rescale_betas_zero_snr:
alphas_cumprod = self.rescale_zero_terminal_snr(alphas_cumprod)
return torch.sqrt(1 - alphas_cumprod)
else:
return torch.sqrt(1 - torch.cumprod(1 - square_betas, dim=0))
def get_schedule(self):
square_betas = torch.linspace(self.beta_min**0.5,
self.beta_max**0.5,
self.num_timesteps,
dtype=torch.float32)**2
self._sigmas = self.square_betas_to_sigmas(square_betas,
self.snr_shift_scale,
self.rescale_betas_zero_snr)
self._betas = torch.sqrt(square_betas)
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)
@NOISE_SCHEDULERS.register_class()
class LinearScheduler(BaseNoiseScheduler):
para_dict = {}
def init_params(self):
super().init_params()
self.beta_min = self.cfg.get('BETA_MIN', 0.00085)
self.beta_max = self.cfg.get('BETA_MAX', 0.012)
def betas_to_sigmas(self, betas):
return torch.sqrt(1 - torch.cumprod(1 - betas, dim=0))
def get_schedule(self):
betas = torch.linspace(self.beta_min,
self.beta_max,
self.num_timesteps,
dtype=torch.float32)
sigmas = self.betas_to_sigmas(betas)
self._sigmas = sigmas
self._betas = betas
self._alphas = torch.sqrt(1 - betas**2)
self._alphas_bar = torch.sqrt(1 - sigmas**2)
self._timesteps = torch.arange(len(sigmas), dtype=torch.float32)
@NOISE_SCHEDULERS.register_class()
class FlowMatchUniformScheduler(BaseNoiseScheduler):
def get_schedule(self):
timesteps = np.linspace(1,
self.num_timesteps,
self.num_timesteps,
dtype=np.float32).copy()
timesteps = torch.from_numpy(timesteps).to(dtype=torch.float32)
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_bar = torch.sqrt(1 - self._sigmas ** 2)
def add_noise(self, x_0, noise=None, t=None, **kwargs):
if t is None:
t = torch.rand(
(x_0.shape[0], ), device=x_0.device) * self.num_timesteps
sigma = self.t_to_sigma(t)
shape = (x_0.size(0), ) + (1, ) * (x_0.ndim - 1)
x_t = (1 - sigma.view(shape)) * x_0 + sigma.view(shape) * noise
return ScheduleOutput(x_0=x_0,
x_t=x_t,
t=t,
sigma=sigma,
alpha_bar=self.t_to_alpha_bar(t))
def sigma_to_t(self, sigma, **kwargs):
return sigma * self.num_timesteps
def t_to_sigma(self, t, **kwargs):
return t / self.num_timesteps
@staticmethod
def get_config_template():
return dict_to_yaml('NOISE_SCHEDULER',
__class__.__name__,
FlowMatchUniformScheduler.para_dict,
set_name=True)
@NOISE_SCHEDULERS.register_class()
class FlowMatchSigmoidScheduler(FlowMatchUniformScheduler):
para_dict = {
'SIGMOID_SCALE': {
'value': 1,
'description': 'The scale for the sigmoid function.'
}
}
def init_params(self):
super().init_params()
self.sigmoid_scale = self.cfg.get('SIGMOID_SCALE', 1)
def sigma_to_t(self, sigma, **kwargs):
t = -torch.log(1 / sigma - 1) / self.sigmoid_scale
return t * self.num_timesteps
def t_to_sigma(self, t, **kwargs):
return torch.sigmoid(self.sigmoid_scale * t / self.num_timesteps)
@staticmethod
def get_config_template():
return dict_to_yaml('NOISE_SCHEDULER',
__class__.__name__,
FlowMatchSigmoidScheduler.para_dict,
set_name=True)
@NOISE_SCHEDULERS.register_class()
class FlowMatchShiftScheduler(FlowMatchUniformScheduler):
para_dict = {
'SHIFT': {
'value': 3,
'description': 'The shift factor for the timestamp.'
},
'SIGMOID_SCALE': {
'value': 1,
'description': 'The scale for the sigmoid function.'
}
}
def init_params(self):
super().init_params()
self.shift = self.cfg.get('SHIFT', 3)
self.sigmoid_scale = self.cfg.get('SIGMOID_SCALE', 1)
def add_noise(self, x_0, noise=None, t=None, **kwargs):
if t is None:
logits_norm = torch.randn(x_0.shape[0], device=x_0.device)
logits_norm = logits_norm * self.sigmoid_scale # larger scale for more uniform sampling
t = logits_norm.sigmoid() * self.num_timesteps
sigma = self.t_to_sigma(t)
shape = (x_0.size(0), ) + (1, ) * (x_0.ndim - 1)
x_t = (1 - sigma.view(shape)) * x_0 + sigma.view(shape) * noise
return ScheduleOutput(x_0=x_0,
x_t=x_t,
t=t,
sigma=sigma,
alpha_bar=self.t_to_alpha_bar(t))
def sigma_to_t(self, sigma, **kwargs):
t = sigma / (sigma - self.shift * sigma + self.shift)
return t * self.num_timesteps
def t_to_sigma(self, t, **kwargs):
t = t / self.num_timesteps
return (t * self.shift) / (1 + (self.shift - 1) * t)
@staticmethod
def get_config_template():
return dict_to_yaml('NOISE_SCHEDULER',
__class__.__name__,
FlowMatchShiftScheduler.para_dict,
set_name=True)
@NOISE_SCHEDULERS.register_class()
class FlowMatchFluxShiftScheduler(FlowMatchUniformScheduler):
para_dict = {
'SHIFT': {
'value': True,
'description': 'Use timestamp shift or not, default is True.'
},
'SIGMOID_SCALE': {
'value': 1,
'description':
'The scale of sigmoid function for sampling timesteps.'
},
'BASE_SHIFT': {
'value': 0.5,
'description': 'The base shift factor for the timestamp.'
},
'MAX_SHIFT': {
'value': 1.15,
'description': 'The max shift factor for the timestamp.'
},
'PRE_T_SAMPLE': {
'value': False,
'description': 'Use pre-sampled timesteps or not, default is False.'
},
'PRE_T_SAMPLE_FOLD': {
'value': 1,
'description': 'The folds of pre-sampled timesteps.'
}
}
def init_params(self):
super().init_params()
self.shift = self.cfg.get('SHIFT', True)
self.sigmoid_scale = self.cfg.get('SIGMOID_SCALE', 1)
self.base_shift = self.cfg.get('BASE_SHIFT', 0.5)
self.max_shift = self.cfg.get('MAX_SHIFT', 1.15)
self.pre_t_sample = self.cfg.get('PRE_T_SAMPLE', False)
self.pre_t_sample_fold = self.cfg.get('PRE_T_SAMPLE_FOLD', 1)
if self.pre_t_sample:
t = torch.sigmoid(torch.randn((self.num_timesteps * self.pre_t_sample_fold,)))
# Scale and reverse the values to go from 1000 to 0
timesteps = ((1 - t) * 1000)
# Sort the timesteps in descending order
self.pre_sample_timesteps, _ = torch.sort(timesteps, descending=True)
else:
self.pre_sample_timesteps = None
@property
def pre_timesteps(self):
fold_id = random.randint(0, self.pre_t_sample_fold - 1)
# print("fold_id", fold_id)
return self.pre_sample_timesteps[fold_id::self.pre_t_sample_fold]
def time_shift(self, mu: float, sigma_scale: float, t: Tensor):
return math.exp(mu) / (math.exp(mu) + (1 / t - 1)**sigma_scale)
def sigma_shift(self, mu: float, sigma_scale: float, sigma: Tensor):
return 1 / (torch.pow(
(1 - sigma) * math.exp(mu) / sigma, sigma_scale) + 1)
def get_lin_function(self,
x1: float = 256,
y1: float = 0.5,
x2: float = 4096,
y2: float = 1.15) -> Callable[[float], float]:
m = (y2 - y1) / (x2 - x1)
b = y1 - m * x1
return lambda x: m * x + b
def add_noise(self, x_0, noise=None, t=None, **kwargs):
if x_0.ndim == 3:
seq_len = x_0.shape[2] // 4
else:
n, _, h, w = x_0.shape
seq_len = (h // 2 * w // 2)
if t is None:
if self.pre_t_sample:
timestep_indices = torch.randint(
1,
self.num_timesteps - 1,
(x_0.shape[0],)
)
timestep_indices = timestep_indices.long()
t = [self.pre_timesteps[x.item()].to(x_0.device) for x in timestep_indices]
t = torch.stack(t, dim=0)
else:
logits_norm = torch.randn(x_0.shape[0], device=x_0.device)
logits_norm = logits_norm * self.sigmoid_scale # larger scale for more uniform sampling
t = logits_norm.sigmoid() * self.num_timesteps
sigma = self.t_to_sigma(t, seq_len=seq_len)
shape = (x_0.size(0), ) + (1, ) * (x_0.ndim - 1)
# print(sigma)
x_t = (1 - sigma.view(shape)) * x_0 + sigma.view(shape) * noise
return ScheduleOutput(x_0=x_0,
x_t=x_t,
t=t,
sigma=sigma,
alpha_bar=self.t_to_alpha_bar(t))
def sigma_to_t(self, sigma, **kwargs):
seq_len = kwargs.get('seq_len', 256)
if self.shift:
mu = self.get_lin_function(y1=self.base_shift,
y2=self.max_shift)(seq_len)
sigma = self.sigma_shift(mu, 1.0, sigma)
t = torch.as_tensor(sigma, dtype=torch.float32)
return t * self.num_timesteps
def t_to_sigma(self, t, **kwargs):
seq_len = kwargs.get('seq_len', 256)
t = t / self.num_timesteps
if self.shift:
mu = self.get_lin_function(y1=self.base_shift,
y2=self.max_shift)(seq_len)
t = self.time_shift(mu, 1.0, t)
sigma = torch.as_tensor(t, dtype=torch.float32)
return sigma
@staticmethod
def get_config_template():
return dict_to_yaml('NOISE_SCHEDULER',
__class__.__name__,
FlowMatchFluxShiftScheduler.para_dict,
set_name=True)
@NOISE_SCHEDULERS.register_class()
class FlowMatchSigmaScheduler(FlowMatchUniformScheduler):
para_dict = {
'WEIGHTING_SCHEME': {
'value':
'logit_normal',
'description':
'The weighting scheme for sampling timesteps, '
"choose from ['sigma_sqrt', 'logit_normal', 'mode', 'cosmap', 'none']."
},
'SHIFT': {
'value': 3.0,
'description': 'The shift factor for the timestamp.'
},
'LOGIT_MEAN': {
'value':
0.0,
'description':
'The mean of the logit distribution for sampling timesteps.'
},
'LOGIT_STD': {
'value':
1.0,
'description':
'The standard deviation of the logit distribution for sampling timesteps.'
},
'MODE_SCALE': {
'value':
1.29,
'description':
'The scale factor for the mode of the logit distribution for sampling timesteps.'
}
}
def init_params(self):
super().init_params()
self.weighting_scheme = self.cfg.get('WEIGHTING_SCHEME',
'logit_normal')
self.logit_mean = self.cfg.get('LOGIT_MEAN', 0.0)
self.logit_std = self.cfg.get('LOGIT_STD', 1.0)
self.mode_scale = self.cfg.get('MODE_SCALE', 1.29)
self.shift = self.cfg.get('SHIFT', 1.0)
def get_schedule(self):
timesteps = np.linspace(1,
self.num_timesteps,
self.num_timesteps,
dtype=np.float32).copy()
timesteps = torch.from_numpy(timesteps).to(dtype=torch.float32)
self._timesteps = timesteps
timesteps = timesteps / self.num_timesteps
self._sigmas = self.shift * timesteps / (1 +
(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:
if self.weighting_scheme == 'logit_normal':
t = torch.normal(mean=self.logit_mean,
std=self.logit_std,
size=(x_0.shape[0], ),
device=x_0.device)
else:
t = torch.rand(x_0.shape[0], device=x_0.device)
t = self.compute_density_for_timestep_sampling(
t) * self.num_timesteps
sigma = self.t_to_sigma(t)
shape = (x_0.size(0), ) + (1, ) * (x_0.ndim - 1)
x_t = (1 - sigma.view(shape)) * x_0 + sigma.view(shape) * noise
return ScheduleOutput(x_0=x_0,
x_t=x_t,
t=t,
sigma=sigma,
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.
Courtesy: This was contributed by Rafie Walker in https://github.com/huggingface/diffusers/pull/8528.
SD3 paper reference: https://arxiv.org/abs/2403.03206v1.
"""
if self.weighting_scheme == 'logit_normal':
# See 3.1 in the SD3 paper ($rf/lognorm(0.00,1.00)$).
t = torch.nn.functional.sigmoid(t)
elif self.weighting_scheme == 'mode':
t = 1 - t - self.mode_scale * (torch.cos(math.pi * t / 2)**2 - 1 +
t)
return t
def sigma_to_t(self, sigma, **kwargs):
raise NotImplementedError
def t_to_sigma(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]
sigma = self.sigmas[step_indices].flatten().to(t)
return sigma
@staticmethod
def get_config_template():
return dict_to_yaml('NOISE_SCHEDULER',
__class__.__name__,
FlowMatchSigmaScheduler.para_dict,
set_name=True)
if __name__ == '__main__':
from scepter.modules.utils.config import Config
cfg = Config(cfg_dict={
'NAME': 'FlowMatchShiftScheduler',
'SHIFT': 1.15
},
load=False)
scheduler = NOISE_SCHEDULERS.build(cfg)