211 lines
9.0 KiB
Python
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) |