Files
modelscope-scepter/scepter/modules/model/diffusion/samplers.py
T
2024-10-21 00:35:53 +08:00

211 lines
9.0 KiB
Python

from dataclasses import dataclass, field
import torch
from scepter.modules.utils.config import dict_to_yaml
from scepter.modules.model.registry import DIFFUSION_SAMPLERS
from .util import _i
@dataclass
class SamplerOutput(object):
callback_fn: callable
prediction_type: str
alphas: torch.Tensor
betas: torch.Tensor
sigmas: torch.Tensor
alphas_init: torch.Tensor
betas_init: torch.Tensor
sigmas_init: torch.Tensor
ts: torch.Tensor
x_t: torch.Tensor
x_0: torch.Tensor
step: int
msg: str
def add_custom_field(self, key: str, value) -> None:
self.__setattr__(key, value)
@DIFFUSION_SAMPLERS.register_class("base")
class BaseDiffusionSampler(object):
para_dict = {}
def __init__(self, cfg, logger=None):
super(BaseDiffusionSampler, self).__init__()
self.logger = logger
self.cfg = cfg
self.init_params()
def init_params(self):
self.discretization_type = self.cfg.get("DISCRETIZATION_TYPE", "linspace")
self.discard_penultimate_step = self.cfg.get("DISCARD_PENULTIMATE_STEP", False)
self.free_steps = self.cfg.get("FREE_STEPS", None)
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):
# get timesteps
if isinstance(steps, int):
steps += 1 if self.discard_penultimate_step else 0
t_max = num_timesteps - 1 if self.t_max is None else self.t_max
t_min = 0 if self.t_min is None else self.t_min
# discretize timesteps
if self.discretization_type == 'leading':
steps = torch.arange(t_min, t_max + 1,
(t_max - t_min + 1) / steps).flip(0)
elif self.discretization_type == 'linspace':
steps = torch.linspace(t_max, t_min, steps)
elif self.discretization_type == 'trailing':
steps = torch.arange(t_max, t_min - 1,
-((t_max - t_min + 1) / steps))
elif self.discretization_type == 'free':
steps = torch.tensor(self.free_steps)
else:
raise NotImplementedError(
f'{self.discretization_type} discretization not implemented')
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
def preprare_sampler(self, noise, steps=20, scheduler_ins=None, prediction_type="",
sigmas=None, betas=None, alphas=None, callback_fn = None,
**kwargs):
'''
1. Control the model's inputs and outputs externally in the solver by callback_fn,
and perform the conversion between x0 and xt internally within the solver.
2. The function callback_fn use the x0 and xt as the default inputs and also give me
the x0 and xt as output. The other inputs will be set in kwargs.
3. The basic parameters of the schedule should be set manually.
4. To ensure the safety of threading, use the instance of SamplerOutput as the manager,
which manage all necessary information.
'''
num_timesteps = scheduler_ins.num_timesteps if scheduler_ins is not None else 1000
timestamps = self.discretization(steps, num_timesteps=num_timesteps, **kwargs)
alphas = scheduler_ins.t_to_alpha(timestamps, **kwargs) if scheduler_ins is not None else alphas
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
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
output = SamplerOutput(
callback_fn=callback_fn,
prediction_type=prediction_type,
alphas=alphas,
betas=betas,
sigmas=sigmas,
alphas_init=alphas_init,
betas_init=betas_init,
sigmas_init=sigmas_init,
ts=timestamps,
x_t=noise,
x_0=noise,
step=0,
msg=f"step 0"
)
return output
def step(self, sampler_ouput):
raise NotImplementedError(f'DiffusionSampler step function not implemented')
def __repr__(self) -> str:
return f'{self.__class__.__name__}' + ' ' + super().__repr__()
@staticmethod
def get_config_template():
return dict_to_yaml('DIFFUSION_SAMPLERS',
__class__.__name__,
BaseDiffusionSampler.para_dict,
set_name=True)
@DIFFUSION_SAMPLERS.register_class("eluer")
class EulerSampler(BaseDiffusionSampler):
def step(self, sampler_ouput):
pass
@DIFFUSION_SAMPLERS.register_class("ddim")
class DDIMSampler(BaseDiffusionSampler):
def init_params(self):
super().init_params()
self.eta = self.cfg.get('ETA', 0.)
self.discretization_type = self.cfg.get("DISCRETIZATION_TYPE", "trailing")
def preprare_sampler(self, noise, steps=20, scheduler_ins=None, prediction_type="",
sigmas=None, betas=None, alphas=None, callback_fn = None,
**kwargs):
output = super().preprare_sampler(noise, steps, scheduler_ins, prediction_type, sigmas, betas, alphas, 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)
return output
def step(self, sampler_output):
step = sampler_output.step
x_t = sampler_output.x_t
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)
sigma_init = _i(sampler_output.sigmas_init, step, x_t)
x = sampler_output.callback_fn(x_t, t, sigma_init, alpha_init)
noise_factor = self.eta * (sigmas_vp[step + 1] ** 2 / sigmas_vp[step] ** 2 *
(1 - (1 - sigmas_vp[step] ** 2) /
(1 - sigmas_vp[step + 1] ** 2)))
d = (x_t - (1 - sigmas_vp[step] ** 2) ** 0.5 * x) / sigmas_vp[step]
x = (1 - sigmas_vp[step + 1] ** 2) ** 0.5 * x + \
(sigmas_vp[step + 1] ** 2 - noise_factor ** 2) ** 0.5 * d
sampler_output.x_0 = x
if sigmas_vp[step + 1] > 0:
x += noise_factor * torch.randn_like(x)
# print("i:", step, "sigma_init:", sigma_init, "alpha_init", alpha_init, "sigmas_vp[i]", sigmas_vp[step], "torch.sum(x_0):", torch.sum(x_0), "torch.sum(x):", torch.sum(x))
sampler_output.x_t = x
sampler_output.step += 1
sampler_output.msg = f'step {step}'
return sampler_output
@DIFFUSION_SAMPLERS.register_class("flow_eluer")
class FlowEluerSampler(BaseDiffusionSampler):
def preprare_sampler(self, noise, steps=20, scheduler_ins=None, prediction_type="",
sigmas=None, betas=None, alphas=None, callback_fn = None,
**kwargs):
if noise.ndim == 3:
seq_len = noise.shape[2] // 4
else:
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)
return output
def step(self, sampler_output):
step = sampler_output.step
x_t = sampler_output.x_t
sigma_curr, sigma_prev = sampler_output.sigmas[step], sampler_output.sigmas[step + 1]
prediction_type = sampler_output.prediction_type
assert prediction_type in ("raw", "sigma_scaled")
t = sampler_output.ts[step]
x_0 = sampler_output.callback_fn(x_t, t, sigma_curr)
x_t = x_t + (sigma_prev - sigma_curr) * x_0
sampler_output.x_0 = x_0
sampler_output.x_t = x_t
sampler_output.step += 1
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):
# extra step for zero
timesteps = torch.linspace(num_timesteps, 0, steps + 1)
return timesteps
@staticmethod
def get_config_template():
return dict_to_yaml('DIFFUSION_SAMPLERS',
__class__.__name__,
FlowEluerSampler.para_dict,
set_name=True)