new project from v0.0.1

This commit is contained in:
duanmurs@163.com
2023-12-28 18:23:39 +08:00
parent fda66e7ca6
commit 66f979f7e9
204 changed files with 37941 additions and 0 deletions
@@ -0,0 +1,7 @@
# -*- coding: utf-8 -*-
# Copyright (c) Alibaba, Inc. and its affiliates.
from scepter.modules.model.network.autoencoder import ae_kl
from scepter.modules.model.network.classifier import Classifier
from scepter.modules.model.network.diffusion import (diffusion, schedules,
solvers)
from scepter.modules.model.network.ldm import ldm, ldm_xl
@@ -0,0 +1,3 @@
# -*- coding: utf-8 -*-
# Copyright (c) Alibaba, Inc. and its affiliates.
from scepter.modules.model.network.autoencoder.ae_kl import AutoencoderKL
@@ -0,0 +1,268 @@
# -*- coding: utf-8 -*-
# Copyright (c) Alibaba, Inc. and its affiliates.
from collections import OrderedDict
import numpy as np
import torch
from scepter.modules.model.network.train_module import TrainModule
from scepter.modules.model.registry import BACKBONES, LOSSES, MODELS
from scepter.modules.utils.config import dict_to_yaml
from scepter.modules.utils.distribute import we
from scepter.modules.utils.file_system import FS
class DiagonalGaussianDistribution(object):
def __init__(self, mean, logvar, deterministic=False):
self.mean = mean
self.logvar = torch.clamp(logvar, -30.0, 20.0)
self.deterministic = deterministic
self.std = torch.exp(0.5 * self.logvar)
self.var = torch.exp(self.logvar)
if self.deterministic:
self.var = self.std = torch.zeros_like(
self.mean).to(device=self.mean.device)
def sample(self):
x = self.mean + self.std * torch.randn(
self.mean.shape).to(device=self.mean.device)
return x
def kl(self, other=None):
if self.deterministic:
return torch.Tensor([0.])
else:
if other is None:
return 0.5 * torch.sum(
torch.pow(self.mean, 2) + self.var - 1.0 - self.logvar,
dim=[1, 2, 3])
else:
return 0.5 * torch.sum(
torch.pow(self.mean - other.mean, 2) / other.var +
self.var / other.var - 1.0 - self.logvar + other.logvar,
dim=[1, 2, 3])
def nll(self, sample, dims=[1, 2, 3]):
if self.deterministic:
return torch.Tensor([0.])
logtwopi = np.log(2.0 * np.pi)
return 0.5 * torch.sum(logtwopi + self.logvar +
torch.pow(sample - self.mean, 2) / self.var,
dim=dims)
def mode(self):
print('*** use DiagonalGaussianDistribution.mode() ***')
return self.mean
@MODELS.register_class()
class AutoencoderKL(TrainModule):
para_dict = {
'ENCODER': {},
'DECODER': {},
'LOSS': {},
'EMBED_DIM': {
'value': 4,
'description': ''
},
'PRETRAINED_MODEL': {
'value': None,
'description': ''
},
'IGNORE_KEYS': {
'value': [],
'description': ''
},
'BATCH_SIZE': {
'value': 16,
'description': ''
},
}
def __init__(self, cfg, logger=None):
super().__init__(cfg, logger=logger)
self.encoder_cfg = self.cfg.ENCODER
self.decoder_cfg = self.cfg.DECODER
self.loss_cfg = self.cfg.get('LOSS', None)
self.embed_dim = self.cfg.get('EMBED_DIM', 4)
self.pretrained_model = self.cfg.get('PRETRAINED_MODEL', None)
self.ignore_keys = self.cfg.get('IGNORE_KEYS', [])
self.batch_size = self.cfg.get('BATCH_SIZE', 16)
self.construct_network()
self.init_network()
def construct_network(self):
z_channels = self.encoder_cfg.Z_CHANNELS
self.encoder = BACKBONES.build(self.encoder_cfg, logger=self.logger)
self.decoder = BACKBONES.build(self.decoder_cfg, logger=self.logger)
self.conv1 = torch.nn.Conv2d(2 * z_channels, 2 * self.embed_dim, 1)
self.conv2 = torch.nn.Conv2d(self.embed_dim, z_channels, 1)
if self.loss_cfg is not None:
self.loss = LOSSES.build(self.loss_cfg, logger=self.logger)
def init_network(self):
if self.pretrained_model is not None:
with FS.get_from(self.pretrained_model,
wait_finish=True) as local_model:
self.init_from_ckpt(local_model, ignore_keys=self.ignore_keys)
def init_from_ckpt(self, path, ignore_keys=list()):
if path.find('.safetensors') > -1:
from safetensors import safe_open
sd = OrderedDict()
with safe_open(path, framework='pt', device='cpu') as f:
for k in f.keys():
sd[k] = f.get_tensor(k)
else:
sd = torch.load(path, map_location='cpu')
if path.find('.pt') > -1 and 'state_dict' in sd:
sd = sd['state_dict']
elif path.find('.ckpt') > -1 and 'state_dict' in sd:
sd = sd['state_dict']
new_sd = OrderedDict()
for k, v in sd.items():
ignored = False
for ik in ignore_keys:
if ik in k:
if we.rank == 0:
self.logger.info(
'ignore key {} from state_dict.'.format(k))
ignored = True
break
k = k.replace('post_quant_conv',
'conv2') if 'post_quant_conv' in k else k
k = k.replace('quant_conv', 'conv1') if 'quant_conv' in k else k
if not ignored:
new_sd[k] = v
missing, unexpected = self.load_state_dict(new_sd, strict=False)
if we.rank == 0:
self.logger.info(
f'Restored from {path} with {len(missing)} missing and {len(unexpected)} unexpected keys'
)
if len(missing) > 0:
self.logger.info(f'Missing Keys:\n {missing}')
if len(unexpected) > 0:
self.logger.info(f'\nUnexpected Keys:\n {unexpected}')
def encode(self, x, return_mom=False):
return torch.cat([
self._encode(batch, return_mom=return_mom)
for batch in x.split(self.batch_size, dim=0)
],
dim=0)
def decode(self, z):
return torch.cat(
[self._decode(batch) for batch in z.split(self.batch_size, dim=0)],
dim=0)
def sample(self, moments):
mean, logvar = torch.chunk(moments, 2, dim=1)
posterior = DiagonalGaussianDistribution(mean, logvar)
z = posterior.sample()
return z
def _encode(self, x, return_mom=False):
h = self.encoder(x)
moments = self.conv1(h)
if return_mom:
return moments
mean, logvar = torch.chunk(moments, 2, dim=1)
posterior = DiagonalGaussianDistribution(mean, logvar)
z = posterior.sample()
return z
def _decode(self, z):
z = self.conv2(z)
dec = self.decoder(z)
return dec
def share_forward(self, image, sample_posterior=True):
posterior = self.encode(image)
if sample_posterior:
z = posterior.sample()
else:
z = posterior.mode()
dec = self.decode(z)
return dec, posterior
def forward(self, **kwargs):
if self.training:
ret = self.forward_train(**kwargs)
else:
ret = self.forward_test(**kwargs)
return ret
def forward_train(self,
image=None,
sample_posterior=True,
optimizer_idx=0,
**kwargs):
reconstructions, posterior = self.share_forward(
image, sample_posterior)
ret = {}
if optimizer_idx == 0:
# train encoder+decoder+logvar
aeloss, log_dict_ae = self.loss(image,
reconstructions,
posterior,
optimizer_idx,
self.global_step,
last_layer=self.get_last_layer(),
split='train')
ret['loss'] = aeloss
ret.update(log_dict_ae)
if optimizer_idx == 1:
# train the discriminator
discloss, log_dict_disc = self.loss(
image,
reconstructions,
posterior,
optimizer_idx,
self.global_step,
last_layer=self.get_last_layer(),
split='train')
ret['loss'] = discloss
ret.update(log_dict_disc)
return ret
def forward_test(self, image=None, sample_posterior=True, **kwargs):
reconstructions, posterior = self.share_forward(
image, sample_posterior)
ret = {}
aeloss, log_dict_ae = self.loss(image,
reconstructions,
posterior,
0,
self.global_step,
last_layer=self.get_last_layer(),
split='val')
discloss, log_dict_disc = self.loss(image,
reconstructions,
posterior,
1,
self.global_step,
last_layer=self.get_last_layer(),
split='val')
ret.update(log_dict_ae)
ret.update(log_dict_disc)
return ret
def get_last_layer(self):
return self.decoder.conv_out.weight
@staticmethod
def get_config_template():
return dict_to_yaml('MODEL',
__class__.__name__,
AutoencoderKL.para_dict,
set_name=True)
+208
View File
@@ -0,0 +1,208 @@
# -*- coding: utf-8 -*-
# Copyright 2021 Alibaba Group Holding Limited. All Rights Reserved.
from collections import OrderedDict
from functools import partial
import torch.nn as nn
from torch.nn.functional import sigmoid, softmax
from scepter.modules.model.metric.registry import METRICS
from scepter.modules.model.network.train_module import TrainModule
from scepter.modules.model.registry import (BACKBONES, HEADS, LOSSES, MODELS,
NECKS)
from scepter.modules.utils.config import Config, dict_to_yaml
_ACTIVATE_MAPPER = {'softmax': partial(softmax, dim=1), 'sigmoid': sigmoid}
@MODELS.register_class()
class Classifier(TrainModule):
""" Base classifier implementation.
Args:
backbones (dict): Defines backbones.
neck (dict, optional): Defines neck. Use Identity if none.
head (dict): Defines head.
act_name (str): Defines activate function, 'softmax' or 'sigmoid'.
topk (Sequence[int]): Defines how to calculate accuracy metrics.
freeze_bn (bool): If True, freeze all BatchNorm layers including LayerNorm.
"""
para_dict = {
'ACT_NAME': {
'value':
'softmax',
'description':
'the activation function for logits, select from [softmax, sigmoid]!'
},
'FREEZE_BN': {
'value': False,
'description': 'if freeze bn of not'
}
}
def __init__(self, cfg, logger=None):
super().__init__(cfg, logger=logger)
# Construct model
self.backbone = BACKBONES.build(cfg.BACKBONE, logger=logger)
necks_cfg = cfg.get('NECK',
Config(cfg_dict={'NAME': 'Identity'}, load=False))
self.neck = NECKS.build(necks_cfg, logger=logger)
self.head = HEADS.build(cfg.HEAD, logger=logger)
freeze_bn = cfg.get('FREEZE_BN', False)
# Construct loss
loss = cfg.get('LOSS',
Config(cfg_dict={'NAME': 'CrossEntropy'}, load=False))
self.loss = LOSSES.build(loss, logger=logger)
act_name = cfg.get('ACT_NAME', 'softmax')
# Construct activate function
self.act_fn = _ACTIVATE_MAPPER[act_name]
self.metric = METRICS.build(cfg.METRIC, logger=logger)
self.freeze_bn = freeze_bn
def train(self, mode=True):
self.training = mode
super(Classifier, self).train(mode=mode)
if self.freeze_bn:
for module in self.modules():
if isinstance(module,
(nn.BatchNorm2d, nn.BatchNorm3d, nn.LayerNorm)):
module.train(False)
return self
def forward(self, img, label=None, **kwargs):
return self.forward_train(
img, label=label) if self.training else self.forward_test(
img, label=label) # noqa
def forward_train(self, img, label=None):
probs = self.head(self.neck(self.backbone(img)))
if label is None:
return probs
ret = OrderedDict()
loss = self.loss(probs, label)
ret['loss'] = loss
ret['batch_size'] = img.size(0)
ret.update(self.metric(probs, label))
return ret
def forward_test(self, img, label=None):
logits = self.act_fn(self.head(self.neck(self.backbone(img))))
if label is not None:
ret = OrderedDict()
ret['logits'] = logits
ret['batch_size'] = img.size(0)
ret.update(self.metric(logits, label))
return ret
return logits
@staticmethod
def get_config_template():
'''
{ "ENV" :
{ "description" : "",
"A" : {
"value": 1.0,
"description": ""
}
}
}
:return:
'''
return dict_to_yaml('MODEL',
__class__.__name__,
Classifier.para_dict,
set_name=True)
@MODELS.register_class()
class VideoClassifier(Classifier):
""" Classifier for video.
Default input tensor is video.
"""
def forward(self, video, label=None, **kwargs):
return self.forward_train(video, label=label) \
if self.training else self.forward_test(video, label=label)
@staticmethod
def get_config_template():
'''
{ "ENV" :
{ "description" : "",
"A" : {
"value": 1.0,
"description": ""
}
}
}
:return:
'''
return dict_to_yaml('MODEL',
__class__.__name__,
VideoClassifier.para_dict,
set_name=True)
@MODELS.register_class()
class VideoClassifier2x(VideoClassifier):
""" A 2-way classifier for video.
"""
def forward_train(self, video, label=None):
probs0, probs1 = self.head(self.neck(self.backbone(video)))
if label is not None:
ret = OrderedDict()
loss = self.loss(probs0, label[:, 0]) + self.loss(
probs1, label[:, 1])
ret['loss'] = loss
ret['batch_size'] = video.size(0)
acc_0 = self.metric(probs0, label[:, 0])
acc_0 = {
key.relace('@', '_0@'): value
for key, value in acc_0.items()
}
acc_1 = self.metric(probs1, label[:, 1])
acc_1 = {
key.relace('@', '_1@'): value
for key, value in acc_1.items()
}
ret.update(acc_0)
ret.update(acc_1)
return ret
return {'logits0': self.act_fn(probs0), 'logits1': self.act_fn(probs1)}
def forward_test(self, video, label=None):
probs0, probs1 = self.head(self.neck(self.backbone(video)))
logits0, logits1 = self.act_fn(probs0), self.act_fn(probs1)
if label is None:
return {'logits0': logits0, 'logits1': logits1}
ret = OrderedDict()
ret['logits0'] = logits0
ret['logits1'] = logits1
acc_0 = self.metric(probs0, label[:, 0])
acc_0 = {key.relace('@', '_0@'): value for key, value in acc_0.items()}
acc_1 = self.metric(probs1, label[:, 1])
acc_1 = {key.relace('@', '_1@'): value for key, value in acc_1.items()}
ret.update(acc_0)
ret.update(acc_1)
return ret
@staticmethod
def get_config_template():
'''
{ "ENV" :
{ "description" : "",
"A" : {
"value": 1.0,
"description": ""
}
}
}
:return:
'''
return dict_to_yaml('MODEL',
__class__.__name__,
VideoClassifier2x.para_dict,
set_name=True)
@@ -0,0 +1,4 @@
# -*- coding: utf-8 -*-
# Copyright (c) Alibaba, Inc. and its affiliates.
from scepter.modules.model.network.diffusion import (diffusion, schedules,
solvers)
@@ -0,0 +1,555 @@
# -*- coding: utf-8 -*-
# Copyright (c) Alibaba, Inc. and its affiliates.
"""
GaussianDiffusion wraps operators for denoising diffusion models, including the
diffusion and denoising processes, as well as the loss evaluation.
"""
import copy
import random
import torch
from .schedules import karras_schedule
from .solvers import (sample_ddim, sample_dpm_2, sample_dpm_2_ancestral,
sample_dpmpp_2m, sample_dpmpp_2m_sde,
sample_dpmpp_2s_ancestral, sample_dpmpp_sde,
sample_euler, sample_euler_ancestral, sample_heun,
sample_img2img_euler, sample_img2img_euler_ancestral)
__all__ = ['GaussianDiffusion']
def _i(tensor, t, x):
"""
Index tensor using t and format the output according to x.
"""
shape = (x.size(0), ) + (1, ) * (x.ndim - 1)
return tensor[t.to(tensor.device)].view(shape).to(x.device)
class GaussianDiffusion(object):
def __init__(self, sigmas, prediction_type='eps'):
assert prediction_type in {'x0', 'eps', 'v'}
self.sigmas = sigmas # noise coefficients
self.alphas = torch.sqrt(1 - sigmas**2) # signal coefficients
self.num_timesteps = len(sigmas)
self.prediction_type = prediction_type
def diffuse(self, x0, t, noise=None):
"""
Add Gaussian noise to signal x0 according to:
q(x_t | x_0) = N(x_t | alpha_t x_0, sigma_t^2 I).
"""
noise = torch.randn_like(x0) if noise is None else noise
xt = _i(self.alphas, t, x0) * x0 + _i(self.sigmas, t, x0) * noise
return xt
def denoise(self,
xt,
t,
s,
model,
model_kwargs={},
guide_scale=None,
guide_rescale=None,
clamp=None,
percentile=None,
cat_uc=False):
"""
Apply one step of denoising from the posterior distribution q(x_s | x_t, x0).
Since x0 is not available, estimate the denoising results using the learned
distribution p(x_s | x_t, \hat{x}_0 == f(x_t)). # noqa
"""
s = t - 1 if s is None else s
# hyperparams
sigmas = _i(self.sigmas, t, xt)
alphas = _i(self.alphas, t, xt)
alphas_s = _i(self.alphas, s.clamp(0), xt)
alphas_s[s < 0] = 1.
sigmas_s = torch.sqrt(1 - alphas_s**2)
# precompute variables
betas = 1 - (alphas / alphas_s)**2
coef1 = betas * alphas_s / sigmas**2
coef2 = (alphas * sigmas_s**2) / (alphas_s * sigmas**2)
var = betas * (sigmas_s / sigmas)**2
log_var = torch.log(var).clamp_(-20, 20)
# prediction
if guide_scale is None:
assert isinstance(model_kwargs, dict)
out = model(xt, t=t, **model_kwargs)
else:
# classifier-free guidance (arXiv:2207.12598)
# model_kwargs[0]: conditional kwargs
# model_kwargs[1]: non-conditional kwargs
assert isinstance(model_kwargs, list) and len(model_kwargs) == 2
if guide_scale == 1.:
out = model(xt, t=t, **model_kwargs[0])
else:
if cat_uc:
def parse_model_kwargs(prev_value, value):
if isinstance(value, torch.Tensor):
prev_value = torch.cat([prev_value, value], dim=0)
elif isinstance(value, dict):
for k, v in value.items():
prev_value[k] = parse_model_kwargs(
prev_value[k], v)
elif isinstance(value, list):
for idx, v in enumerate(value):
prev_value[idx] = parse_model_kwargs(
prev_value[idx], v)
return prev_value
all_model_kwargs = copy.deepcopy(model_kwargs[0])
for model_kwarg in model_kwargs[1:]:
for key, value in model_kwarg.items():
all_model_kwargs[key] = parse_model_kwargs(
all_model_kwargs[key], value)
all_out = model(xt.repeat(2, 1, 1, 1),
t=t.repeat(2),
**all_model_kwargs)
y_out, u_out = all_out.chunk(2)
else:
y_out = model(xt, t=t, **model_kwargs[0])
u_out = model(xt, t=t, **model_kwargs[1])
out = u_out + guide_scale * (y_out - u_out)
# rescale the output according to arXiv:2305.08891
if guide_rescale is not None:
assert guide_rescale >= 0 and guide_rescale <= 1
ratio = (y_out.flatten(1).std(dim=1) /
(out.flatten(1).std(dim=1) +
1e-12)).view((-1, ) + (1, ) * (y_out.ndim - 1))
out *= guide_rescale * ratio + (1 - guide_rescale) * 1.0
# compute x0
if self.prediction_type == 'x0':
x0 = out
elif self.prediction_type == 'eps':
x0 = (xt - sigmas * out) / alphas
elif self.prediction_type == 'v':
x0 = alphas * xt - sigmas * out
else:
raise NotImplementedError(
f'prediction_type {self.prediction_type} not implemented')
# restrict the range of x0
if percentile is not None:
# NOTE: percentile should only be used when data is within range [-1, 1]
assert percentile > 0 and percentile <= 1
s = torch.quantile(x0.flatten(1).abs(), percentile, dim=1)
s = s.clamp_(1.0).view((-1, ) + (1, ) * (xt.ndim - 1))
x0 = torch.min(s, torch.max(-s, x0)) / s
elif clamp is not None:
x0 = x0.clamp(-clamp, clamp)
# recompute eps using the restricted x0
eps = (xt - alphas * x0) / sigmas
# compute mu (mean of posterior distribution) using the restricted x0
mu = coef1 * x0 + coef2 * xt
return mu, var, log_var, x0, eps
def loss(self,
x0,
t,
model,
model_kwargs={},
reduction='mean',
noise=None):
# hyperparams
sigmas = _i(self.sigmas, t, x0)
alphas = _i(self.alphas, t, x0)
# diffuse and denoise
if noise is None:
noise = torch.randn_like(x0)
xt = self.diffuse(x0, t, noise)
out = model(xt, t=t, **model_kwargs)
# mse loss
target = {
'eps': noise,
'x0': x0,
'v': alphas * noise - sigmas * x0
}[self.prediction_type]
loss = (out - target).pow(2)
if reduction == 'mean':
loss = loss.flatten(1).mean(dim=1)
return loss
@torch.no_grad()
def sample(self,
noise,
model,
x=None,
denoising_strength=1.0,
refine_stage=False,
refine_strength=0.0,
model_kwargs={},
condition_fn=None,
guide_scale=None,
guide_rescale=None,
clamp=None,
percentile=None,
solver='euler_a',
steps=20,
t_max=None,
t_min=None,
discretization=None,
discard_penultimate_step=None,
return_intermediate=None,
show_progress=False,
seed=-1,
intermediate_callback=None,
cat_uc=False,
**kwargs):
# sanity check
assert isinstance(steps, (int, torch.LongTensor))
assert t_max is None or (t_max > 0 and t_max <= self.num_timesteps - 1)
assert t_min is None or (t_min >= 0 and t_min < self.num_timesteps - 1)
assert discretization in (None, 'leading', 'linspace', 'trailing')
assert discard_penultimate_step in (None, True, False)
assert return_intermediate in (None, 'x0', 'xt')
# function of diffusion solver
solver_fn = {
'ddim': sample_ddim,
'euler_ancestral': sample_euler_ancestral,
'euler': sample_euler,
'heun': sample_heun,
'dpm2': sample_dpm_2,
'dpm2_ancestral': sample_dpm_2_ancestral,
'dpmpp_2s_ancestral': sample_dpmpp_2s_ancestral,
'dpmpp_2m': sample_dpmpp_2m,
'dpmpp_sde': sample_dpmpp_sde,
'dpmpp_2m_sde': sample_dpmpp_2m_sde,
'dpm2_karras': sample_dpm_2,
'dpm2_ancestral_karras': sample_dpm_2_ancestral,
'dpmpp_2s_ancestral_karras': sample_dpmpp_2s_ancestral,
'dpmpp_2m_karras': sample_dpmpp_2m,
'dpmpp_sde_karras': sample_dpmpp_sde,
'dpmpp_2m_sde_karras': sample_dpmpp_2m_sde
}[solver]
# options
schedule = 'karras' if 'karras' in solver else None
discretization = discretization or 'linspace'
seed = seed if seed >= 0 else random.randint(0, 2**31)
if isinstance(steps, torch.LongTensor):
discard_penultimate_step = False
if discard_penultimate_step is None:
discard_penultimate_step = True if solver in (
'dpm2', 'dpm2_ancestral', 'dpmpp_2m_sde', 'dpm2_karras',
'dpm2_ancestral_karras', 'dpmpp_2m_sde_karras') else False
# function for denoising xt to get x0
intermediates = []
def model_fn(xt, sigma):
# denoising
t = self._sigma_to_t(sigma).repeat(len(xt)).round().long()
x0 = self.denoise(xt,
t,
None,
model,
model_kwargs,
guide_scale,
guide_rescale,
clamp,
percentile,
cat_uc=cat_uc)[-2]
# collect intermediate outputs
if return_intermediate == 'xt':
intermediates.append(xt)
elif return_intermediate == 'x0':
intermediates.append(x0)
if intermediate_callback is not None:
intermediate_callback(intermediates[-1])
return x0
# get timesteps
if isinstance(steps, int):
steps += 1 if discard_penultimate_step else 0
t_max = self.num_timesteps - 1 if t_max is None else t_max
t_min = 0 if t_min is None else t_min
# discretize timesteps
if discretization == 'leading':
steps = torch.arange(t_min, t_max + 1,
(t_max - t_min + 1) / steps).flip(0)
elif discretization == 'linspace':
steps = torch.linspace(t_max, t_min, steps)
elif discretization == 'trailing':
steps = torch.arange(t_max, t_min - 1,
-((t_max - t_min + 1) / steps))
else:
raise NotImplementedError(
f'{discretization} discretization not implemented')
steps = steps.clamp_(t_min, t_max)
steps = torch.as_tensor(steps,
dtype=torch.float32,
device=noise.device)
# get sigmas
sigmas = self._t_to_sigma(steps)
sigmas = torch.cat([sigmas, sigmas.new_zeros([1])])
t_enc = int(min(denoising_strength, 0.999) * len(steps))
sigmas = sigmas[len(steps) - t_enc - 1:]
if refine_strength > 0:
t_refine = int(min(refine_strength, 0.999) * len(steps))
if refine_stage:
sigmas = sigmas[-t_refine:]
else:
sigmas = sigmas[:-t_refine + 1]
# print(sigmas)
if x is not None:
noise = (x + noise * sigmas[0]) / torch.sqrt(1.0 + sigmas[0]**2.0)
if schedule == 'karras':
if sigmas[0] == float('inf'):
sigmas = karras_schedule(
n=len(steps) - 1,
sigma_min=sigmas[sigmas > 0].min().item(),
sigma_max=sigmas[sigmas < float('inf')].max().item(),
rho=7.).to(sigmas)
sigmas = torch.cat([
sigmas.new_tensor([float('inf')]), sigmas,
sigmas.new_zeros([1])
])
else:
sigmas = karras_schedule(
n=len(steps),
sigma_min=sigmas[sigmas > 0].min().item(),
sigma_max=sigmas.max().item(),
rho=7.).to(sigmas)
sigmas = torch.cat([sigmas, sigmas.new_zeros([1])])
if discard_penultimate_step:
sigmas = torch.cat([sigmas[:-2], sigmas[-1:]])
kwargs['seed'] = seed
# sampling
x0 = solver_fn(noise,
model_fn,
sigmas,
show_progress=show_progress,
**kwargs)
return (x0, intermediates) if return_intermediate is not None else x0
def _sigma_to_t(self, sigma):
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):
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()
@torch.no_grad()
def stochastic_encode(self, x0, t, steps):
# fast, but does not allow for exact reconstruction
# t serves as an index to gather the correct alphas
t_max = None
t_min = None
# discretization method
discretization = 'trailing' if self.prediction_type == 'v' else 'leading'
# timesteps
if isinstance(steps, int):
t_max = self.num_timesteps - 1 if t_max is None else t_max
t_min = 0 if t_min is None else t_min
steps = discretize_timesteps(t_max, t_min, steps, discretization)
steps = torch.as_tensor(steps).round().long().flip(0).to(x0.device)
# steps = torch.as_tensor(steps).round().long().to(x0.device)
# self.alphas_bar = torch.cumprod(1 - self.sigmas ** 2, dim=0)
# print('sigma: ', self.sigmas, len(self.sigmas))
# print('alpha_bar: ', self.alphas_bar, len(self.alphas_bar))
# print('steps: ', steps, len(steps))
# sqrt_alphas_cumprod = torch.sqrt(self.alphas_bar).to(x0.device)[steps]
# sqrt_one_minus_alphas_cumprod = torch.sqrt(1 - self.alphas_bar).to(x0.device)[steps]
sqrt_alphas_cumprod = self.alphas.to(x0.device)[steps]
sqrt_one_minus_alphas_cumprod = self.sigmas.to(x0.device)[steps]
# print('sigma: ', self.sigmas, len(self.sigmas))
# print('alpha: ', self.alphas, len(self.alphas))
# print('steps: ', steps, len(steps))
noise = torch.randn_like(x0)
return (
extract_into_tensor(sqrt_alphas_cumprod, t, x0.shape) * x0 +
extract_into_tensor(sqrt_one_minus_alphas_cumprod, t, x0.shape) *
noise)
@torch.no_grad()
def sample_img2img(self,
x,
noise,
model,
denoising_strength=1,
model_kwargs={},
condition_fn=None,
guide_scale=None,
guide_rescale=None,
clamp=None,
percentile=None,
solver='euler_a',
steps=20,
t_max=None,
t_min=None,
discretization=None,
discard_penultimate_step=None,
return_intermediate=None,
show_progress=False,
seed=-1,
**kwargs):
# sanity check
assert isinstance(steps, (int, torch.LongTensor))
assert t_max is None or (t_max > 0 and t_max <= self.num_timesteps - 1)
assert t_min is None or (t_min >= 0 and t_min < self.num_timesteps - 1)
assert discretization in (None, 'leading', 'linspace', 'trailing')
assert discard_penultimate_step in (None, True, False)
assert return_intermediate in (None, 'x0', 'xt')
# function of diffusion solver
solver_fn = {
'euler_ancestral': sample_img2img_euler_ancestral,
'euler': sample_img2img_euler,
}[solver]
# options
schedule = 'karras' if 'karras' in solver else None
discretization = discretization or 'linspace'
seed = seed if seed >= 0 else random.randint(0, 2**31)
if isinstance(steps, torch.LongTensor):
discard_penultimate_step = False
if discard_penultimate_step is None:
discard_penultimate_step = True if solver in (
'dpm2', 'dpm2_ancestral', 'dpmpp_2m_sde', 'dpm2_karras',
'dpm2_ancestral_karras', 'dpmpp_2m_sde_karras') else False
# function for denoising xt to get x0
intermediates = []
def get_scalings(sigma):
c_out = -sigma
c_in = 1 / (sigma**2 + 1.**2)**0.5
return c_out, c_in
def model_fn(xt, sigma):
# denoising
c_out, c_in = get_scalings(sigma)
t = self._sigma_to_t(sigma).repeat(len(xt)).round().long()
x0 = self.denoise(xt * c_in, t, None, model, model_kwargs,
guide_scale, guide_rescale, clamp,
percentile)[-2]
# collect intermediate outputs
if return_intermediate == 'xt':
intermediates.append(xt)
elif return_intermediate == 'x0':
intermediates.append(x0)
return xt + x0 * c_out
# get timesteps
if isinstance(steps, int):
steps += 1 if discard_penultimate_step else 0
t_max = self.num_timesteps - 1 if t_max is None else t_max
t_min = 0 if t_min is None else t_min
# discretize timesteps
if discretization == 'leading':
steps = torch.arange(t_min, t_max + 1,
(t_max - t_min + 1) / steps).flip(0)
elif discretization == 'linspace':
steps = torch.linspace(t_max, t_min, steps)
elif discretization == 'trailing':
steps = torch.arange(t_max, t_min - 1,
-((t_max - t_min + 1) / steps))
else:
raise NotImplementedError(
f'{discretization} discretization not implemented')
steps = steps.clamp_(t_min, t_max)
steps = torch.as_tensor(steps, dtype=torch.float32, device=x.device)
# get sigmas
sigmas = self._t_to_sigma(steps)
sigmas = torch.cat([sigmas, sigmas.new_zeros([1])])
t_enc = int(min(denoising_strength, 0.999) * len(steps))
sigmas = sigmas[len(steps) - t_enc - 1:]
noise = x + noise * sigmas[0]
if schedule == 'karras':
if sigmas[0] == float('inf'):
sigmas = karras_schedule(
n=len(steps) - 1,
sigma_min=sigmas[sigmas > 0].min().item(),
sigma_max=sigmas[sigmas < float('inf')].max().item(),
rho=7.).to(sigmas)
sigmas = torch.cat([
sigmas.new_tensor([float('inf')]), sigmas,
sigmas.new_zeros([1])
])
else:
sigmas = karras_schedule(
n=len(steps),
sigma_min=sigmas[sigmas > 0].min().item(),
sigma_max=sigmas.max().item(),
rho=7.).to(sigmas)
sigmas = torch.cat([sigmas, sigmas.new_zeros([1])])
if discard_penultimate_step:
sigmas = torch.cat([sigmas[:-2], sigmas[-1:]])
# sampling
x0 = solver_fn(noise,
model_fn,
sigmas,
seed=seed,
show_progress=show_progress,
**kwargs)
return (x0, intermediates) if return_intermediate is not None else x0
def extract_into_tensor(a, t, x_shape):
b, *_ = t.shape
out = a.gather(-1, t)
return out.reshape(b, *((1, ) * (len(x_shape) - 1)))
def discretize_timesteps(t_max, t_min, steps, discretization):
"""
Implementation of timestep discretization methods.
"""
if discretization == 'leading':
steps = torch.arange(t_min, t_max + 1,
(t_max - t_min + 1) / steps).flip(0)
elif discretization == 'linspace':
steps = torch.linspace(t_max, t_min, steps)
elif discretization == 'trailing':
steps = torch.arange(t_max, t_min - 1, -((t_max - t_min + 1) / steps))
else:
raise NotImplementedError(
f'{discretization} discretization not implemented')
return steps.clamp_(t_min, t_max)
@@ -0,0 +1,181 @@
# -*- coding: utf-8 -*-
# Copyright (c) Alibaba, Inc. and its affiliates.
"""
Noise schedules of denoising diffusion probabilistic models.
We consider a variance preserving (VP) process, and we use the standard deviation
sigma_t of the noise added to the signal at time t to represent the noise schedule. The
corresponding diffusion process is:
q(x_t | x_0) = N(x_t | alpha_t x_0, sigma_t^2 I),
where alpha_t^2 = 1 - sigma_t^2.
"""
import math
import torch
__all__ = [
'betas_to_sigmas', 'sigmas_to_betas', 'logsnrs_to_sigmas',
'sigmas_to_logsnrs', 'linear_schedule', 'quadratic_schedule',
'scaled_linear_schedule', 'cosine_schedule', 'sigmoid_schedule',
'karras_schedule', 'exponential_schedule', 'polyexponential_schedule',
'vp_schedule', 'logsnr_cosine_schedule', 'logsnr_cosine_shifted_schedule',
'logsnr_cosine_interp_schedule', 'noise_schedule'
]
def betas_to_sigmas(betas):
return torch.sqrt(1 - torch.cumprod(1 - betas, dim=0))
def sigmas_to_betas(sigmas):
square_alphas = 1 - sigmas**2
betas = 1 - torch.cat(
[square_alphas[:1], square_alphas[1:] / square_alphas[:-1]])
return betas
def logsnrs_to_sigmas(logsnrs):
return torch.sqrt(torch.sigmoid(-logsnrs))
def sigmas_to_logsnrs(sigmas):
square_sigmas = sigmas**2
return torch.log(square_sigmas / (1 - square_sigmas))
def linear_schedule(n, beta_min=0.0001, beta_max=0.02):
betas = torch.linspace(beta_min, beta_max, n, dtype=torch.float32)
return betas_to_sigmas(betas)
def scaled_linear_schedule(n, beta_min=0.00085, beta_max=0.012):
betas = torch.linspace(beta_min**0.5,
beta_max**0.5,
n,
dtype=torch.float32)**2
return betas_to_sigmas(betas)
def quadratic_schedule(n=1000, init_beta=0.00085, last_beta=0.012):
betas = torch.linspace(init_beta**0.5,
last_beta**0.5,
n,
dtype=torch.float32)**2
return betas_to_sigmas(betas)
def cosine_schedule(n, cosine_s=0.008):
ramp = torch.linspace(0, 1, n + 1)
square_alphas = torch.cos(
(ramp + cosine_s) / (1 + cosine_s) * torch.pi / 2)**2
betas = (1 - square_alphas[1:] / square_alphas[:-1]).clamp(max=0.999)
return betas_to_sigmas(betas)
def sigmoid_schedule(n, beta_min=0.0001, beta_max=0.02):
betas = torch.sigmoid(torch.linspace(-6, 6,
n)) * (beta_max - beta_min) + beta_min
return betas_to_sigmas(betas)
def karras_schedule(n, sigma_min=0.002, sigma_max=80.0, rho=7.0):
ramp = torch.linspace(1, 0, n)
min_inv_rho = sigma_min**(1 / rho)
max_inv_rho = sigma_max**(1 / rho)
sigmas = (max_inv_rho + ramp * (min_inv_rho - max_inv_rho))**rho
sigmas = torch.sqrt(sigmas**2 / (1 + sigmas**2)) # VE -> VP
return sigmas
def exponential_schedule(n, sigma_min=0.002, sigma_max=80.0):
sigmas = torch.linspace(math.log(sigma_min), math.log(sigma_max), n).exp()
sigmas = torch.sqrt(sigmas**2 / (1 + sigmas**2)) # VE -> VP
return sigmas
def polyexponential_schedule(n, sigma_min=0.002, sigma_max=80.0):
ramp = torch.linspace(0, 1, n)
sigmas = torch.exp(ramp * (math.log(sigma_max) - math.log(sigma_min)) +
math.log(sigma_min))
sigmas = torch.sqrt(sigmas**2 / (1 + sigmas**2)) # VE -> VP
return sigmas
def vp_schedule(n, beta_d=19.9, beta_min=0.1, eps_s=1e-3):
t = torch.linspace(eps_s, 1, n)
sigmas = torch.sqrt(torch.exp(beta_d * t**2 / 2 + beta_min * t) - 1)
sigmas = torch.sqrt(sigmas**2 / (1 + sigmas**2)) # VE -> VP
return sigmas
def _logsnr_cosine(n, logsnr_min=-15, logsnr_max=15):
t_min = math.atan(math.exp(-0.5 * logsnr_min))
t_max = math.atan(math.exp(-0.5 * logsnr_max))
t = torch.linspace(1, 0, n)
logsnrs = -2 * torch.log(torch.tan(t_min + t * (t_max - t_min)))
return logsnrs
def _logsnr_cosine_shifted(n, logsnr_min=-15, logsnr_max=15, scale=2):
logsnrs = _logsnr_cosine(n, logsnr_min, logsnr_max)
logsnrs += 2 * math.log(1 / scale)
return logsnrs
def _logsnr_cosine_interp(n,
logsnr_min=-15,
logsnr_max=15,
scale_min=2,
scale_max=4):
t = torch.linspace(1, 0, n)
logsnrs_min = _logsnr_cosine_shifted(n, logsnr_min, logsnr_max, scale_min)
logsnrs_max = _logsnr_cosine_shifted(n, logsnr_min, logsnr_max, scale_max)
logsnrs = t * logsnrs_min + (1 - t) * logsnrs_max
return logsnrs
def logsnr_cosine_schedule(n, logsnr_min=-15, logsnr_max=15):
return logsnrs_to_sigmas(_logsnr_cosine(n, logsnr_min, logsnr_max))
def logsnr_cosine_shifted_schedule(n, logsnr_min=-15, logsnr_max=15, scale=2):
return logsnrs_to_sigmas(
_logsnr_cosine_shifted(n, logsnr_min, logsnr_max, scale))
def logsnr_cosine_interp_schedule(n,
logsnr_min=-15,
logsnr_max=15,
scale_min=2,
scale_max=4):
return logsnrs_to_sigmas(
_logsnr_cosine_interp(n, logsnr_min, logsnr_max, scale_min, scale_max))
def noise_schedule(schedule='logsnr_cosine_interp',
n=1000,
zero_terminal_snr=False,
**kwargs):
# compute sigmas
sigmas = {
'linear': linear_schedule,
'scaled_linear': scaled_linear_schedule,
'quadratic': quadratic_schedule,
'cosine': cosine_schedule,
'sigmoid': sigmoid_schedule,
'karras': karras_schedule,
'exponential': exponential_schedule,
'polyexponential': polyexponential_schedule,
'vp': vp_schedule,
'logsnr_cosine': logsnr_cosine_schedule,
'logsnr_cosine_shifted': logsnr_cosine_shifted_schedule,
'logsnr_cosine_interp': logsnr_cosine_interp_schedule
}[schedule](n, **kwargs)
# post-processing
if zero_terminal_snr and sigmas.max() != 1.0:
scale = (1.0 - sigmas.min()) / (sigmas.max() - sigmas.min())
sigmas = sigmas.min() + scale * (sigmas - sigmas.min())
return sigmas
@@ -0,0 +1,611 @@
# -*- coding: utf-8 -*-
# Copyright (c) Alibaba, Inc. and its affiliates.
"""
ODE/SDE solver for denoising diffusion models under either variation preserving (VP) or
variation exploding (VE) settings. Under the VE setting, the diffusion process is:
q(x_t | x_0) = N(x_t | x_0, sigma_t^2 I),
where 0 <= sigma_t <= inf; while under the VP setting, the diffusion process is:
q(x_t | x_0) = N(x_t | alpha_t x_0, sigma_t^2 I),
where 0 <= sigma_t <= 1 and alpha_t^2 = 1 - sigma_t^2.
"""
import torch
from tqdm.auto import trange
__all__ = [
'sample_euler', 'sample_euler_ancestral', 'sample_heun', 'sample_dpm_2',
'sample_dpm_2_ancestral', 'sample_dpmpp_2s_ancestral', 'sample_dpmpp_sde',
'sample_dpmpp_2m', 'sample_dpmpp_2m_sde', 'sample_ddim'
]
# -------------------- variation exploding (VE) solver --------------------#
def get_ancestral_step(sigma_from, sigma_to, eta=1.):
"""
Calculates the noise level (sigma_down) to step down to and the amount
of noise to add (sigma_up) when doing an ancestral sampling step.
"""
if not eta:
return sigma_to, 0.
sigma_up = min(
sigma_to,
eta * (sigma_to**2 *
(sigma_from**2 - sigma_to**2) / sigma_from**2)**0.5)
sigma_down = (sigma_to**2 - sigma_up**2)**0.5
return sigma_down, sigma_up
def get_scalings(sigma):
c_out = -sigma
c_in = 1 / (sigma**2 + 1.**2)**0.5
return c_out, c_in
@torch.no_grad()
def sample_euler(noise,
model,
sigmas,
s_churn=0.,
s_tmin=0.,
s_tmax=float('inf'),
s_noise=1.,
seed=None,
show_progress=True):
"""
Implements Algorithm 2 (Euler steps) from Karras et al. (2022).
"""
x = noise * sigmas[0]
for i in trange(len(sigmas) - 1, disable=not show_progress):
gamma = 0.
if s_tmin <= sigmas[i] <= s_tmax and sigmas[i] < float('inf'):
gamma = min(s_churn / (len(sigmas) - 1), 2**0.5 - 1)
eps = torch.randn_like(x) * s_noise
sigma_hat = sigmas[i] * (gamma + 1)
if gamma > 0:
x = x + eps * (sigma_hat**2 - sigmas[i]**2)**0.5
# Euler method
if sigmas[i] == float('inf'):
denoised = model(noise, sigma_hat)
x = denoised + sigmas[i + 1] * (gamma + 1) * noise
else:
_, c_in = get_scalings(sigma_hat)
denoised = model(x * c_in, sigma_hat)
d = (x - denoised) / sigma_hat
dt = sigmas[i + 1] - sigma_hat
x = x + d * dt
return x
@torch.no_grad()
def sample_euler_ancestral(noise,
model,
sigmas,
eta=1.,
s_noise=1.,
seed=None,
show_progress=True):
"""
Ancestral sampling with Euler method steps.
"""
x = noise * sigmas[0]
for i in trange(len(sigmas) - 1, disable=not show_progress):
sigma_down, sigma_up = get_ancestral_step(sigmas[i],
sigmas[i + 1],
eta=eta)
# Euler method
if sigmas[i] == float('inf'):
denoised = model(noise, sigmas[i])
x = denoised + sigmas[i + 1] * noise
else:
_, c_in = get_scalings(sigmas[i])
denoised = model(x * c_in, sigmas[i])
d = (x - denoised) / sigmas[i]
dt = sigma_down - sigmas[i]
x = x + d * dt
if sigmas[i + 1] > 0:
x = x + torch.randn_like(x) * s_noise * sigma_up
return x
@torch.no_grad()
def sample_heun(noise,
model,
sigmas,
s_churn=0.,
s_tmin=0.,
s_tmax=float('inf'),
s_noise=1.,
seed=None,
show_progress=True):
"""
Implements Algorithm 2 (Heun steps) from Karras et al. (2022).
"""
x = noise * sigmas[0]
for i in trange(len(sigmas) - 1, disable=not show_progress):
gamma = 0.
if s_tmin <= sigmas[i] <= s_tmax and sigmas[i] < float('inf'):
gamma = min(s_churn / (len(sigmas) - 1), 2**0.5 - 1)
eps = torch.randn_like(x) * s_noise
sigma_hat = sigmas[i] * (gamma + 1)
if gamma > 0:
x = x + eps * (sigma_hat**2 - sigmas[i]**2)**0.5
if sigmas[i] == float('inf'):
# Euler method
denoised = model(noise, sigma_hat)
x = denoised + sigmas[i + 1] * (gamma + 1) * noise
else:
_, c_in = get_scalings(sigma_hat)
denoised = model(x * c_in, sigma_hat)
d = (x - denoised) / sigma_hat
dt = sigmas[i + 1] - sigma_hat
if sigmas[i + 1] == 0:
# Euler method
x = x + d * dt
else:
# Heun's method
x_2 = x + d * dt
_, c_in = get_scalings(sigmas[i + 1])
denoised_2 = model(x_2 * c_in, sigmas[i + 1])
d_2 = (x_2 - denoised_2) / sigmas[i + 1]
d_prime = (d + d_2) / 2
x = x + d_prime * dt
return x
@torch.no_grad()
def sample_dpm_2(noise,
model,
sigmas,
s_churn=0.,
s_tmin=0.,
s_tmax=float('inf'),
s_noise=1.,
seed=None,
show_progress=True):
"""
A sampler inspired by DPM-Solver-2 and Algorithm 2 from Karras et al. (2022).
"""
x = noise * sigmas[0]
for i in trange(len(sigmas) - 1, disable=not show_progress):
gamma = 0.
if s_tmin <= sigmas[i] <= s_tmax and sigmas[i] < float('inf'):
gamma = min(s_churn / (len(sigmas) - 1), 2**0.5 - 1)
eps = torch.randn_like(x) * s_noise
sigma_hat = sigmas[i] * (gamma + 1)
if gamma > 0:
x = x + eps * (sigma_hat**2 - sigmas[i]**2)**0.5
if sigmas[i] == float('inf'):
# Euler method
denoised = model(noise, sigma_hat)
x = denoised + sigmas[i + 1] * (gamma + 1) * noise
else:
_, c_in = get_scalings(sigma_hat)
denoised = model(x * c_in, sigma_hat)
d = (x - denoised) / sigma_hat
if sigmas[i + 1] == 0:
# Euler method
dt = sigmas[i + 1] - sigma_hat
x = x + d * dt
else:
# DPM-Solver-2
sigma_mid = sigma_hat.log().lerp(sigmas[i + 1].log(),
0.5).exp()
dt_1 = sigma_mid - sigma_hat
dt_2 = sigmas[i + 1] - sigma_hat
x_2 = x + d * dt_1
_, c_in = get_scalings(sigma_mid)
denoised_2 = model(x_2 * c_in, sigma_mid)
d_2 = (x_2 - denoised_2) / sigma_mid
x = x + d_2 * dt_2
return x
@torch.no_grad()
def sample_dpm_2_ancestral(noise,
model,
sigmas,
eta=1.,
s_noise=1.,
seed=None,
show_progress=True):
"""
Ancestral sampling with DPM-Solver second-order steps.
"""
x = noise * sigmas[0]
for i in trange(len(sigmas) - 1, disable=not show_progress):
sigma_down, sigma_up = get_ancestral_step(sigmas[i],
sigmas[i + 1],
eta=eta)
if sigmas[i] == float('inf'):
# Euler method
denoised = model(noise, sigmas[i])
x = denoised + sigmas[i + 1] * noise
else:
_, c_in = get_scalings(sigmas[i])
denoised = model(x * c_in, sigmas[i])
d = (x - denoised) / sigmas[i]
if sigma_down == 0:
# Euler method
dt = sigma_down - sigmas[i]
x = x + d * dt
else:
# DPM-Solver-2
sigma_mid = sigmas[i].log().lerp(sigma_down.log(), 0.5).exp()
dt_1 = sigma_mid - sigmas[i]
dt_2 = sigma_down - sigmas[i]
x_2 = x + d * dt_1
_, c_in = get_scalings(sigma_mid)
denoised_2 = model(x_2 * c_in, sigma_mid)
d_2 = (x_2 - denoised_2) / sigma_mid
x = x + d_2 * dt_2
x = x + torch.randn_like(x) * s_noise * sigma_up
return x
@torch.no_grad()
def sample_dpmpp_2s_ancestral(noise,
model,
sigmas,
eta=1.,
s_noise=1.,
seed=None,
show_progress=True):
"""
Ancestral sampling with DPM-Solver++ (2S) second-order steps.
"""
def t_to_sigma(t):
return t.neg().exp()
def sigma_to_t(sigma):
return sigma.log().neg()
# x = noise * sigmas[0]
x = noise * torch.sqrt(1.0 + sigmas[0]**2.0)
for i in trange(len(sigmas) - 1, disable=not show_progress):
sigma_down, sigma_up = get_ancestral_step(sigmas[i],
sigmas[i + 1],
eta=eta)
if sigmas[i] == float('inf'):
# Euler method
denoised = model(noise, sigmas[i])
x = denoised + sigma_down * noise
else:
_, c_in = get_scalings(sigmas[i])
denoised = model(x * c_in, sigmas[i])
if sigma_down == 0:
# Euler method
d = (x - denoised) / sigmas[i]
dt = sigma_down - sigmas[i]
x = x + d * dt
else:
# DPM-Solver++(2S)
t, t_next = sigma_to_t(sigmas[i]), sigma_to_t(sigma_down)
r = 1 / 2
h = t_next - t
s = t + r * h
x_2 = (t_to_sigma(s) /
t_to_sigma(t)) * x - (-h * r).expm1() * denoised
_, c_in = get_scalings(t_to_sigma(s))
denoised_2 = model(x_2 * c_in, t_to_sigma(s))
x = (t_to_sigma(t_next) /
t_to_sigma(t)) * x - (-h).expm1() * denoised_2
# Noise addition
if sigmas[i + 1] > 0:
x = x + torch.randn_like(x) * s_noise * sigma_up
return x
class BatchedBrownianTree:
"""
A wrapper around torchsde.BrownianTree that enables batches of entropy.
"""
def __init__(self, x, t0, t1, seed=None, **kwargs):
import torchsde
t0, t1, self.sign = self.sort(t0, t1)
w0 = kwargs.get('w0', torch.zeros_like(x))
if seed is None:
seed = torch.randint(0, 2**63 - 1, []).item()
self.batched = True
try:
assert len(seed) == x.shape[0]
w0 = w0[0]
except TypeError:
seed = [seed]
self.batched = False
self.trees = [
torchsde.BrownianTree(t0, w0, t1, entropy=s, **kwargs)
for s in seed
]
@staticmethod
def sort(a, b):
return (a, b, 1) if a < b else (b, a, -1)
def __call__(self, t0, t1):
t0, t1, sign = self.sort(t0, t1)
w = torch.stack([tree(t0, t1)
for tree in self.trees]) * (self.sign * sign)
return w if self.batched else w[0]
class BrownianTreeNoiseSampler:
"""
A noise sampler backed by a torchsde.BrownianTree.
Args:
x (Tensor): The tensor whose shape, device and dtype to use to generate
random samples.
sigma_min (float): The low end of the valid interval.
sigma_max (float): The high end of the valid interval.
seed (int or List[int]): The random seed. If a list of seeds is
supplied instead of a single integer, then the noise sampler will
use one BrownianTree per batch item, each with its own seed.
transform (callable): A function that maps sigma to the sampler's
internal timestep.
"""
def __init__(self,
x,
sigma_min,
sigma_max,
seed=None,
transform=lambda x: x):
self.transform = transform
t0 = self.transform(torch.as_tensor(sigma_min))
t1 = self.transform(torch.as_tensor(sigma_max))
self.tree = BatchedBrownianTree(x, t0, t1, seed)
def __call__(self, sigma, sigma_next):
t0 = self.transform(torch.as_tensor(sigma))
t1 = self.transform(torch.as_tensor(sigma_next))
return self.tree(t0, t1) / (t1 - t0).abs().sqrt()
@torch.no_grad()
def sample_dpmpp_sde(noise,
model,
sigmas,
eta=1.,
s_noise=1.,
r=1 / 2,
seed=None,
show_progress=True):
"""
DPM-Solver++ (stochastic).
"""
def t_to_sigma(t):
return t.neg().exp()
def sigma_to_t(sigma):
return sigma.log().neg()
x = noise * sigmas[0]
sigma_min, sigma_max = sigmas[sigmas > 0].min(), sigmas[
sigmas < float('inf')].max()
noise_sampler = BrownianTreeNoiseSampler(x, sigma_min, sigma_max, seed)
for i in trange(len(sigmas) - 1, disable=not show_progress):
if sigmas[i] == float('inf'):
# Euler method
denoised = model(noise, sigmas[i])
x = denoised + sigmas[i + 1] * noise
else:
_, c_in = get_scalings(sigmas[i])
denoised = model(x * c_in, sigmas[i])
if sigmas[i + 1] == 0:
# Euler method
d = (x - denoised) / sigmas[i]
dt = sigmas[i + 1] - sigmas[i]
x = x + d * dt
else:
# DPM-Solver++
t, t_next = sigma_to_t(sigmas[i]), sigma_to_t(sigmas[i + 1])
h = t_next - t
s = t + h * r
fac = 1 / (2 * r)
# Step 1
sd, su = get_ancestral_step(t_to_sigma(t), t_to_sigma(s), eta)
s_ = sigma_to_t(sd)
x_2 = (t_to_sigma(s_) /
t_to_sigma(t)) * x - (t - s_).expm1() * denoised
x_2 = x_2 + noise_sampler(t_to_sigma(t),
t_to_sigma(s)) * s_noise * su
_, c_in = get_scalings(t_to_sigma(s))
denoised_2 = model(x_2 * c_in, t_to_sigma(s))
# Step 2
sd, su = get_ancestral_step(t_to_sigma(t), t_to_sigma(t_next),
eta)
t_next_ = sigma_to_t(sd)
denoised_d = (1 - fac) * denoised + fac * denoised_2
x = (t_to_sigma(t_next_) / t_to_sigma(t)) * x - \
(t - t_next_).expm1() * denoised_d
x = x + noise_sampler(t_to_sigma(t),
t_to_sigma(t_next)) * s_noise * su
return x
@torch.no_grad()
def sample_dpmpp_2m(noise, model, sigmas, seed=None, show_progress=True):
"""
DPM-Solver++ (2M).
"""
def t_to_sigma(t):
return t.neg().exp()
def sigma_to_t(sigma):
return sigma.log().neg()
x = noise * sigmas[0]
old_denoised = None
for i in trange(len(sigmas) - 1, disable=not show_progress):
if sigmas[i] == float('inf'):
# Euler method
denoised = model(noise, sigmas[i])
x = denoised + sigmas[i + 1] * noise
else:
_, c_in = get_scalings(sigmas[i])
denoised = model(x * c_in, sigmas[i])
t, t_next = sigma_to_t(sigmas[i]), sigma_to_t(sigmas[i + 1])
h = t_next - t
if (old_denoised is None or sigmas[i - 1] == float('inf')
or sigmas[i + 1] == 0):
x = (t_to_sigma(t_next) /
t_to_sigma(t)) * x - (-h).expm1() * denoised
else:
h_last = t - sigma_to_t(sigmas[i - 1])
r = h_last / h
denoised_d = (1 + 1 /
(2 * r)) * denoised - (1 /
(2 * r)) * old_denoised
x = (t_to_sigma(t_next) /
t_to_sigma(t)) * x - (-h).expm1() * denoised_d
old_denoised = denoised
return x
@torch.no_grad()
def sample_dpmpp_2m_sde(noise,
model,
sigmas,
eta=1.,
s_noise=1.,
solver_type='midpoint',
seed=None,
show_progress=True):
"""
DPM-Solver++ (2M) SDE.
"""
assert solver_type in {'heun', 'midpoint'}
x = noise * sigmas[0]
sigma_min, sigma_max = sigmas[sigmas > 0].min(), sigmas[
sigmas < float('inf')].max()
noise_sampler = BrownianTreeNoiseSampler(x, sigma_min, sigma_max, seed)
old_denoised = None
h_last = None
for i in trange(len(sigmas) - 1, disable=not show_progress):
if sigmas[i] == float('inf'):
# Euler method
denoised = model(noise, sigmas[i])
x = denoised + sigmas[i + 1] * noise
else:
_, c_in = get_scalings(sigmas[i])
denoised = model(x * c_in, sigmas[i])
if sigmas[i + 1] == 0:
# Denoising step
x = denoised
else:
# DPM-Solver++(2M) SDE
t, s = -sigmas[i].log(), -sigmas[i + 1].log()
h = s - t
eta_h = eta * h
x = sigmas[i + 1] / sigmas[i] * (-eta_h).exp() * x + \
(-h - eta_h).expm1().neg() * denoised
if old_denoised is not None:
r = h_last / h
if solver_type == 'heun':
x = x + ((-h - eta_h).expm1().neg() / (-h - eta_h) + 1) * \
(1 / r) * (denoised - old_denoised)
elif solver_type == 'midpoint':
x = x + 0.5 * (-h - eta_h).expm1().neg() * \
(1 / r) * (denoised - old_denoised)
x = x + noise_sampler(sigmas[i], sigmas[i + 1]) * sigmas[
i + 1] * (-2 * eta_h).expm1().neg().sqrt() * s_noise
old_denoised = denoised
h_last = h
return x
# -------------------- variation preserving (VP) solver --------------------#
@torch.no_grad()
def sample_ddim(noise, model, sigmas, eta=0., seed=None, show_progress=True):
"""
DDIM solver steps.
"""
x = noise
sigmas_vp = (sigmas**2 / (1 + sigmas**2))**0.5
sigmas_vp[sigmas == float('inf')] = 1.
for i in trange(len(sigmas) - 1, disable=not show_progress):
denoised = model(x, sigmas[i])
noise_factor = eta * (sigmas_vp[i + 1]**2 / sigmas_vp[i]**2 *
(1 - (1 - sigmas_vp[i]**2) /
(1 - sigmas_vp[i + 1]**2)))
d = (x - (1 - sigmas_vp[i]**2)**0.5 * denoised) / sigmas_vp[i]
x = (1 - sigmas_vp[i + 1] ** 2) ** 0.5 * denoised + \
(sigmas_vp[i + 1] ** 2 - noise_factor ** 2) ** 0.5 * d
if sigmas_vp[i + 1] > 0:
x += noise_factor * torch.randn_like(x)
return x
@torch.no_grad()
def sample_img2img_euler(noise,
model,
sigmas,
s_churn=0.,
s_tmin=0.,
s_tmax=float('inf'),
s_noise=1.,
seed=None,
show_progress=True):
"""
Implements Algorithm 2 (Euler steps) from Karras et al. (2022).
"""
x = noise
for i in trange(len(sigmas) - 1, disable=not show_progress):
gamma = 0.
if s_tmin <= sigmas[i] <= s_tmax and sigmas[i] < float('inf'):
gamma = min(s_churn / (len(sigmas) - 1), 2**0.5 - 1)
eps = torch.randn_like(x) * s_noise
sigma_hat = sigmas[i] * (gamma + 1)
if gamma > 0:
x = x + eps * (sigma_hat**2 - sigmas[i]**2)**0.5
# Euler method
if sigmas[i] == float('inf'):
denoised = model(noise, sigma_hat)
x = denoised + sigmas[i + 1] * (gamma + 1) * noise
else:
denoised = model(x, sigma_hat)
d = (x - denoised) / sigma_hat
dt = sigmas[i + 1] - sigma_hat
x = x + d * dt
return x
@torch.no_grad()
def sample_img2img_euler_ancestral(noise,
model,
sigmas,
eta=1.,
s_noise=1.,
seed=None,
show_progress=True):
"""
Ancestral sampling with Euler method steps.
"""
x = noise
for i in trange(len(sigmas) - 1, disable=not show_progress):
sigma_down, sigma_up = get_ancestral_step(sigmas[i],
sigmas[i + 1],
eta=eta)
# Euler method
if sigmas[i] == float('inf'):
denoised = model(noise, sigmas[i])
x = denoised + sigmas[i + 1] * noise
else:
denoised = model(x, sigmas[i])
d = (x - denoised) / sigmas[i]
dt = sigma_down - sigmas[i]
x = x + d * dt
if sigmas[i + 1] > 0:
x = x + torch.randn_like(x) * s_noise * sigma_up
return x
@@ -0,0 +1,4 @@
# -*- coding: utf-8 -*-
# Copyright (c) Alibaba, Inc. and its affiliates.
from scepter.modules.model.network.ldm.ldm import LatentDiffusion
from scepter.modules.model.network.ldm.ldm_xl import LatentDiffusionXL
+432
View File
@@ -0,0 +1,432 @@
# -*- coding: utf-8 -*-
# Copyright (c) Alibaba, Inc. and its affiliates.
import numbers
import random
from collections import OrderedDict
import torch
from scepter.modules.model.network.diffusion.diffusion import GaussianDiffusion
from scepter.modules.model.network.diffusion.schedules import noise_schedule
from scepter.modules.model.network.train_module import TrainModule
from scepter.modules.model.registry import (BACKBONES, EMBEDDERS, LOSSES,
MODELS, TOKENIZERS)
from scepter.modules.model.utils.basic_utils import count_params, default
from scepter.modules.utils.config import dict_to_yaml
from scepter.modules.utils.distribute import we
from scepter.modules.utils.file_system import FS
def disabled_train(self, mode=True):
"""Overwrite model.train with this function to make sure train/eval mode
does not change anymore."""
return self
@MODELS.register_class()
class LatentDiffusion(TrainModule):
para_dict = {
'PARAMETERIZATION': {
'value':
'v',
'description':
"The prediction type, you can choose from 'eps' and 'x0' and 'v'",
},
'TIMESTEPS': {
'value': 1000,
'description': 'The schedule steps for diffusion.',
},
'SCHEDULE_ARGS': {},
'MIN_SNR_GAMMA': {
'value': None,
'description': 'The minimum snr gamma, default is None.',
},
'ZERO_TERMINAL_SNR': {
'value': False,
'description': 'Whether zero terminal snr, default is False.',
},
'PRETRAINED_MODEL': {
'value': None,
'description': "Whole model's pretrained model path.",
},
'IGNORE_KEYS': {
'value': [],
'description': 'The ignore keys for pretrain model loaded.',
},
'SCALE_FACTOR': {
'value': 0.18215,
'description': 'The vae embeding scale.',
},
'SIZE_FACTOR': {
'value': 8,
'description': 'The vae size factor.',
},
'DEFAULT_N_PROMPT': {
'value': '',
'description': 'The default negtive prompt.',
},
'TRAIN_N_PROMPT': {
'value': '',
'description': 'The negtive prompt used in train phase.',
},
'P_ZERO': {
'value': 0.0,
'description': 'The prob for zero or negtive prompt.',
},
'USE_EMA': {
'value': True,
'description': 'Use Ema or not. Default True',
},
'DIFFUSION_MODEL': {},
'DIFFUSION_MODEL_EMA': {},
'FIRST_STAGE_MODEL': {},
'COND_STAGE_MODEL': {},
'TOKENIZER': {}
}
def __init__(self, cfg, logger=None):
super().__init__(cfg, logger=logger)
self.init_params()
self.construct_network()
def init_params(self):
self.parameterization = self.cfg.get('PARAMETERIZATION', 'eps')
assert self.parameterization in [
'eps', 'x0', 'v'
], 'currently only supporting "eps" and "x0" and "v"'
self.num_timesteps = self.cfg.get('TIMESTEPS', 1000)
self.schedule_args = {
k.lower(): v
for k, v in self.cfg.get('SCHEDULE_ARGS', {
'NAME': 'logsnr_cosine_interp',
'SCALE_MIN': 2.0,
'SCALE_MAX': 4.0
}).items()
}
self.min_snr_gamma = self.cfg.get('MIN_SNR_GAMMA', None)
self.zero_terminal_snr = self.cfg.get('ZERO_TERMINAL_SNR', False)
if self.zero_terminal_snr:
assert self.parameterization == 'v', 'Now zero_terminal_snr only support v-prediction mode.'
self.sigmas = noise_schedule(schedule=self.schedule_args.pop('name'),
n=self.num_timesteps,
zero_terminal_snr=self.zero_terminal_snr,
**self.schedule_args)
self.diffusion = GaussianDiffusion(
sigmas=self.sigmas, prediction_type=self.parameterization)
self.pretrained_model = self.cfg.get('PRETRAINED_MODEL', None)
self.ignore_keys = self.cfg.get('IGNORE_KEYS', [])
self.model_config = self.cfg.DIFFUSION_MODEL
self.first_stage_config = self.cfg.FIRST_STAGE_MODEL
self.cond_stage_config = self.cfg.COND_STAGE_MODEL
self.tokenizer_config = self.cfg.get('TOKENIZER', None)
self.loss_config = self.cfg.get('LOSS', None)
self.scale_factor = self.cfg.get('SCALE_FACTOR', 0.18215)
self.size_factor = self.cfg.get('SIZE_FACTOR', 8)
self.default_n_prompt = self.cfg.get('DEFAULT_N_PROMPT', '')
self.default_n_prompt = '' if self.default_n_prompt is None else self.default_n_prompt
self.p_zero = self.cfg.get('P_ZERO', 0.0)
self.train_n_prompt = self.cfg.get('TRAIN_N_PROMPT', '')
if self.default_n_prompt is None:
self.default_n_prompt = ''
if self.train_n_prompt is None:
self.train_n_prompt = ''
self.use_ema = self.cfg.get('USE_EMA', True)
self.model_ema_config = self.cfg.get('DIFFUSION_MODEL_EMA', None)
def construct_network(self):
self.model = BACKBONES.build(self.model_config, logger=self.logger)
self.logger.info('all parameters:{}'.format(count_params(self.model)))
if self.use_ema and self.model_ema_config:
self.model_ema = BACKBONES.build(self.model_ema_config,
logger=self.logger)
self.model_ema = self.model_ema.eval()
for param in self.model_ema.parameters():
param.requires_grad = False
if self.loss_config:
self.loss = LOSSES.build(self.loss_config, logger=self.logger)
if self.tokenizer_config is not None:
self.tokenizer = TOKENIZERS.build(self.tokenizer_config,
logger=self.logger)
self.first_stage_model = MODELS.build(self.first_stage_config,
logger=self.logger)
self.first_stage_model = self.first_stage_model.eval()
self.first_stage_model.train = disabled_train
for param in self.first_stage_model.parameters():
param.requires_grad = False
if self.tokenizer_config is not None:
self.cond_stage_config.KWARGS = {
'vocab_size': self.tokenizer.vocab_size
}
if self.cond_stage_config == '__is_unconditional__':
print(
f'Training {self.__class__.__name__} as an unconditional model.'
)
self.cond_stage_model = None
else:
model = EMBEDDERS.build(self.cond_stage_config, logger=self.logger)
self.cond_stage_model = model.eval().requires_grad_(False)
self.cond_stage_model.train = disabled_train
def load_pretrained_model(self, pretrained_model):
if pretrained_model is not None:
with FS.get_from(pretrained_model,
wait_finish=True) as local_model:
self.init_from_ckpt(local_model, ignore_keys=self.ignore_keys)
def init_from_ckpt(self, path, ignore_keys=list()):
if path.endswith('safetensors'):
from safetensors.torch import load_file as load_safetensors
sd = load_safetensors(path)
else:
sd = torch.load(path, map_location='cpu')
new_sd = OrderedDict()
for k, v in sd.items():
ignored = False
for ik in ignore_keys:
if ik in k:
if we.rank == 0:
self.logger.info(
'Ignore key {} from state_dict.'.format(k))
ignored = True
break
if not ignored:
if k.startswith('model.diffusion_model.'):
k = k.replace('model.diffusion_model.', 'model.')
k = k.replace('post_quant_conv',
'conv2') if 'post_quant_conv' in k else k
k = k.replace('quant_conv',
'conv1') if 'quant_conv' in k else k
new_sd[k] = v
missing, unexpected = self.load_state_dict(new_sd, strict=False)
if we.rank == 0:
self.logger.info(
f'Restored from {path} with {len(missing)} missing and {len(unexpected)} unexpected keys'
)
if len(missing) > 0:
self.logger.info(f'Missing Keys:\n {missing}')
if len(unexpected) > 0:
self.logger.info(f'\nUnexpected Keys:\n {unexpected}')
def encode_condition(self, input, method='encode_text'):
if hasattr(self.cond_stage_model, method):
return getattr(self.cond_stage_model,
method)(input, tokenizer=self.tokenizer)
else:
return self.cond_stage_model(input)
def forward_train(self, image=None, noise=None, prompt=None, **kwargs):
x_start = self.encode_first_stage(image, **kwargs)
t = torch.randint(0,
self.num_timesteps, (x_start.shape[0], ),
device=x_start.device).long()
context = {}
if prompt and self.cond_stage_model:
zeros = (torch.rand(len(prompt)) < self.p_zero).numpy().tolist()
prompt = [
self.train_n_prompt if zeros[idx] else p
for idx, p in enumerate(prompt)
]
self.register_probe({'after_prompt': prompt})
with torch.autocast(device_type='cuda', enabled=False):
context = self.encode_condition(
self.tokenizer(prompt).to(we.device_id))
if self.min_snr_gamma is not None:
alphas = self.diffusion.alphas.to(we.device_id)[t]
sigmas = self.diffusion.sigmas.pow(2).to(we.device_id)[t]
snrs = (alphas / sigmas).clamp(min=1e-20)
min_snrs = snrs.clamp(max=self.min_snr_gamma)
weights = min_snrs / snrs
else:
weights = 1
self.register_probe({'snrs_weights': weights})
loss = self.diffusion.loss(x0=x_start,
t=t,
model=self.model,
model_kwargs={'cond': context},
noise=noise)
loss = loss * weights
loss = loss.mean()
ret = {'loss': loss, 'probe_data': {'prompt': prompt}}
return ret
def noise_sample(self, batch_size, h, w, g):
noise = torch.empty(batch_size, 4, h, w,
device=we.device_id).normal_(generator=g)
return noise
def forward(self, **kwargs):
if self.training:
return self.forward_train(**kwargs)
else:
return self.forward_test(**kwargs)
@torch.no_grad()
@torch.autocast('cuda', dtype=torch.float16)
def forward_test(self,
prompt=None,
n_prompt=None,
sampler='ddim',
sample_steps=50,
seed=2023,
guide_scale=7.5,
guide_rescale=0.5,
discretization='trailing',
run_train_n=True,
**kwargs):
g = torch.Generator(device=we.device_id)
seed = seed if seed >= 0 else random.randint(0, 2**32 - 1)
g.manual_seed(seed)
num_samples = len(prompt)
if 'dynamic_encode_text' in kwargs and kwargs.pop(
'dynamic_encode_text'):
method = 'dynamic_encode_text'
else:
method = 'encode_text'
n_prompt = default(n_prompt, [self.default_n_prompt] * len(prompt))
assert isinstance(prompt, list) and \
isinstance(n_prompt, list) and \
len(prompt) == len(n_prompt)
# with torch.autocast(device_type="cuda", enabled=False):
context = self.encode_condition(self.tokenizer(prompt).to(
we.device_id),
method=method)
null_context = self.encode_condition(self.tokenizer(n_prompt).to(
we.device_id),
method=method)
if 'index' in kwargs:
kwargs.pop('index')
image_size = None
if 'meta' in kwargs:
meta = kwargs.pop('meta')
if 'image_size' in meta:
h = int(meta['image_size'][0][0])
w = int(meta['image_size'][1][0])
image_size = [h, w]
if 'image_size' in kwargs:
image_size = kwargs.pop('image_size')
if image_size is None or isinstance(image_size, numbers.Number):
image_size = [1024, 1024]
height, width = image_size
noise = self.noise_sample(num_samples, height // self.size_factor,
width // self.size_factor, g)
# UNet use input n_prompt
samples = self.diffusion.sample(solver=sampler,
noise=noise,
model=self.model,
model_kwargs=[{
'cond': context
}, {
'cond': null_context
}],
steps=sample_steps,
guide_scale=guide_scale,
guide_rescale=guide_rescale,
discretization=discretization,
show_progress=True,
seed=seed,
condition_fn=None,
clamp=None,
percentile=None,
t_max=None,
t_min=None,
discard_penultimate_step=None,
return_intermediate=None,
**kwargs)
x_samples = self.decode_first_stage(samples).float()
x_samples = torch.clamp((x_samples + 1.0) / 2.0, min=0.0, max=1.0)
# UNet use train n_prompt
if not self.default_n_prompt == self.train_n_prompt and run_train_n:
train_n_prompt = [self.train_n_prompt] * len(prompt)
null_train_context = self.encode_condition(
self.tokenizer(train_n_prompt).to(we.device_id), method=method)
tn_samples = self.diffusion.sample(solver=sampler,
noise=noise,
model=self.model,
model_kwargs=[{
'cond': context
}, {
'cond':
null_train_context
}],
steps=sample_steps,
guide_scale=guide_scale,
guide_rescale=guide_rescale,
discretization=discretization,
show_progress=we.rank == 0,
seed=seed,
condition_fn=None,
clamp=None,
percentile=None,
t_max=None,
t_min=None,
discard_penultimate_step=None,
return_intermediate=None,
**kwargs)
t_x_samples = self.decode_first_stage(tn_samples).float()
t_x_samples = torch.clamp((t_x_samples + 1.0) / 2.0,
min=0.0,
max=1.0)
else:
train_n_prompt = ['' for _ in prompt]
t_x_samples = [None for _ in prompt]
outputs = list()
for p, np, tnp, img, t_img in zip(prompt, n_prompt, train_n_prompt,
x_samples, t_x_samples):
one_tup = {'prompt': p, 'n_prompt': np, 'image': img}
if t_img is not None:
one_tup['train_n_prompt'] = tnp
one_tup['train_n_image'] = t_img
outputs.append(one_tup)
return outputs
@torch.no_grad()
def log_images(self, image=None, prompt=None, n_prompt=None, **kwargs):
results = self.forward_test(prompt=prompt, n_prompt=n_prompt, **kwargs)
outputs = list()
for img, res in zip(image, results):
one_tup = {
'orig': torch.clamp((img + 1.0) / 2.0, min=0.0, max=1.0),
'recon': res['image'],
'prompt': res['prompt'],
'n_prompt': res['n_prompt']
}
if 'train_n_prompt' in res:
one_tup['train_n_prompt'] = res['train_n_prompt']
one_tup['train_n_image'] = res['train_n_image']
outputs.append(one_tup)
return outputs
@torch.no_grad()
def encode_first_stage(self, x, **kwargs):
z = self.first_stage_model.encode(x)
return self.scale_factor * z
@torch.no_grad()
def decode_first_stage(self, z):
z = 1. / self.scale_factor * z
return self.first_stage_model.decode(z)
@staticmethod
def get_config_template():
return dict_to_yaml('MODEL',
__class__.__name__,
LatentDiffusion.para_dict,
set_name=True)
+507
View File
@@ -0,0 +1,507 @@
# -*- coding: utf-8 -*-
# Copyright (c) Alibaba, Inc. and its affiliates.
import numbers
import random
from collections import OrderedDict
import torch
import torch.nn.functional as F
from scepter.modules.model.network.ldm import LatentDiffusion
from scepter.modules.model.registry import BACKBONES, MODELS
from scepter.modules.model.utils.basic_utils import default
from scepter.modules.utils.config import dict_to_yaml
from scepter.modules.utils.distribute import we
@MODELS.register_class()
class LatentDiffusionXL(LatentDiffusion):
para_dict = {
'LOAD_REFINER': {
'value': False,
'description': 'Whether load REFINER or Not.'
}
}
def __init__(self, cfg, logger=None):
super().__init__(cfg, logger=logger)
self.load_refiner = cfg.get('LOAD_REFINER', False)
self.latent_cache_data = {}
self.SURPPORT_RATIOS = {
'0.5': (704, 1408),
'0.52': (704, 1344),
'0.57': (768, 1344),
'0.6': (768, 1280),
'0.68': (832, 1216),
'0.72': (832, 1152),
'0.78': (896, 1152),
'0.82': (896, 1088),
'0.88': (960, 1088),
'0.94': (960, 1024),
'1.0': (1024, 1024),
'1.07': (1024, 960),
'1.13': (1088, 960),
'1.21': (1088, 896),
'1.29': (1152, 896),
'1.38': (1152, 832),
'1.46': (1216, 832),
'1.67': (1280, 768),
'1.75': (1344, 768),
'1.91': (1344, 704),
'2.0': (1408, 704),
'2.09': (1472, 704),
'2.4': (1536, 640),
'2.5': (1600, 640),
'2.89': (1664, 576),
'3.0': (1728, 576),
}
def construct_network(self):
super().construct_network()
self.refiner_cfg = self.cfg.get('REFINER_MODEL', None)
self.refiner_cond_cfg = self.cfg.get('REFINER_COND_MODEL', None)
if self.refiner_cfg and self.load_refiner:
self.refiner_model = BACKBONES.build(self.refiner_cfg,
logger=self.logger)
self.refiner_cond_model = BACKBONES.build(self.refiner_cond_cfg,
logger=self.logger)
else:
self.refiner_model = None
self.refiner_cond_model = None
self.input_keys = self.get_unique_embedder_keys_from_conditioner(
self.cond_stage_model)
if self.refiner_cond_model:
self.input_refiner_keys = self.get_unique_embedder_keys_from_conditioner(
self.refiner_cond_model)
else:
self.input_refiner_keys = []
def init_from_ckpt(self, path, ignore_keys=list()):
if path.endswith('safetensors'):
from safetensors.torch import load_file as load_safetensors
sd = load_safetensors(path)
else:
sd = torch.load(path, map_location='cpu')
new_sd = OrderedDict()
for k, v in sd.items():
ignored = False
for ik in ignore_keys:
if ik in k:
if we.rank == 0:
self.logger.info(
'Ignore key {} from state_dict.'.format(k))
ignored = True
break
if not ignored:
if k.startswith('model.diffusion_model.'):
k = k.replace('model.diffusion_model.', 'model.')
if k.startswith('conditioner.'):
k = k.replace('conditioner.', 'cond_stage_model.')
k = k.replace('post_quant_conv',
'conv2') if 'post_quant_conv' in k else k
k = k.replace('quant_conv',
'conv1') if 'quant_conv' in k else k
new_sd[k] = v
missing, unexpected = self.load_state_dict(new_sd, strict=False)
if we.rank == 0:
self.logger.info(
f'Restored from {path} with {len(missing)} missing and {len(unexpected)} unexpected keys'
)
if len(missing) > 0:
self.logger.info(f'Missing Keys:\n {missing}')
if len(unexpected) > 0:
self.logger.info(f'\nUnexpected Keys:\n {unexpected}')
def get_unique_embedder_keys_from_conditioner(self, conditioner):
input_keys = []
for x in conditioner.embedders:
input_keys.extend(x.input_keys)
return list(set(input_keys))
def get_batch(self, keys, value_dict, num_samples=1):
batch = {}
batch_uc = {}
N = num_samples
device = we.device_id
for key in keys:
if key == 'prompt':
batch['prompt'] = value_dict['prompt']
batch_uc['prompt'] = value_dict['negative_prompt']
elif key == 'original_size_as_tuple':
batch['original_size_as_tuple'] = (torch.tensor(
value_dict['original_size_as_tuple']).to(device).repeat(
N, 1))
elif key == 'crop_coords_top_left':
batch['crop_coords_top_left'] = (torch.tensor(
value_dict['crop_coords_top_left']).to(device).repeat(
N, 1))
elif key == 'aesthetic_score':
batch['aesthetic_score'] = (torch.tensor(
[value_dict['aesthetic_score']]).to(device).repeat(N, 1))
batch_uc['aesthetic_score'] = (torch.tensor([
value_dict['negative_aesthetic_score']
]).to(device).repeat(N, 1))
elif key == 'target_size_as_tuple':
batch['target_size_as_tuple'] = (torch.tensor(
value_dict['target_size_as_tuple']).to(device).repeat(
N, 1))
else:
batch[key] = value_dict[key]
for key in batch.keys():
if key not in batch_uc and isinstance(batch[key], torch.Tensor):
batch_uc[key] = torch.clone(batch[key])
return batch, batch_uc
def forward_train(self, image=None, noise=None, prompt=None, **kwargs):
with torch.autocast('cuda', enabled=False):
x_start = self.encode_first_stage(image, **kwargs)
t = torch.randint(0,
self.num_timesteps, (x_start.shape[0], ),
device=x_start.device).long()
if prompt and self.cond_stage_model:
zeros = (torch.rand(len(prompt)) < self.p_zero).numpy().tolist()
prompt = [
self.train_n_prompt if zeros[idx] else p
for idx, p in enumerate(prompt)
]
self.register_probe({'after_prompt': prompt})
batch = {'prompt': prompt}
for key in self.input_keys:
if key not in kwargs:
continue
batch[key] = kwargs[key].to(we.device_id)
context = getattr(self.cond_stage_model, 'encode')(batch)
if self.min_snr_gamma is not None:
alphas = self.diffusion.alphas.to(we.device_id)[t]
sigmas = self.diffusion.sigmas.pow(2).to(we.device_id)[t]
snrs = (alphas / sigmas).clamp(min=1e-20)
min_snrs = snrs.clamp(max=self.min_snr_gamma)
weights = min_snrs / snrs
else:
weights = 1
self.register_probe({'snrs_weights': weights})
loss = self.diffusion.loss(x0=x_start,
t=t,
model=self.model,
model_kwargs={'cond': context},
noise=noise)
loss = loss * weights
loss = loss.mean()
ret = {'loss': loss, 'probe_data': {'prompt': prompt}}
return ret
def check_valid_inputs(self, kwargs):
batch_data = {}
all_keys = set(self.input_keys + self.input_refiner_keys)
for key in all_keys:
if key in kwargs:
batch_data[key] = kwargs.pop(key)
return batch_data
@torch.no_grad()
def forward_test(self,
prompt=None,
n_prompt=None,
image=None,
sampler='ddim',
sample_steps=50,
seed=2023,
guide_scale=7.5,
guide_rescale=0.5,
discretization='trailing',
img_to_img_strength=0.0,
run_train_n=True,
refine_strength=0.0,
refine_sampler='ddim',
**kwargs):
g = torch.Generator(device=we.device_id)
seed = seed if seed >= 0 else random.randint(0, 2**32 - 1)
g.manual_seed(seed)
num_samples = len(prompt)
n_prompt = default(n_prompt, [self.default_n_prompt] * len(prompt))
assert isinstance(prompt, list) and \
isinstance(n_prompt, list) and \
len(prompt) == len(n_prompt)
image_size = None
if 'meta' in kwargs:
meta = kwargs.pop('meta')
if 'image_size' in meta:
h = int(meta['image_size'][0][0])
w = int(meta['image_size'][1][0])
image_size = [h, w]
if 'image_size' in kwargs:
image_size = kwargs.pop('image_size')
if image_size is None or isinstance(image_size, numbers.Number):
image_size = [1024, 1024]
pre_batch = self.check_valid_inputs(kwargs)
if len(pre_batch) > 0:
batch = {'prompt': prompt}
batch.update(pre_batch)
batch_uc = {'prompt': n_prompt}
batch_uc.update(pre_batch)
else:
height, width = image_size
if image is None:
ori_width = width
ori_height = height
else:
ori_height, ori_width = image.shape[-2:]
value_dict = {
'original_size_as_tuple': [ori_height, ori_width],
'target_size_as_tuple': [height, width],
'prompt': prompt,
'negative_prompt': n_prompt,
'crop_coords_top_left': [0, 0]
}
if refine_strength > 0:
assert 'aesthetic_score' in kwargs and 'negative_aesthetic_score' in kwargs
value_dict['aesthetic_score'] = kwargs.pop('aesthetic_score')
value_dict['negative_aesthetic_score'] = kwargs.pop(
'negative_aesthetic_score')
batch, batch_uc = self.get_batch(self.input_keys,
value_dict,
num_samples=num_samples)
context = getattr(self.cond_stage_model, 'encode')(batch)
null_context = getattr(self.cond_stage_model, 'encode')(batch_uc)
if 'index' in kwargs:
kwargs.pop('index')
height, width = batch['target_size_as_tuple'][0].cpu().numpy().tolist()
noise = self.noise_sample(num_samples, height // self.size_factor,
width // self.size_factor, g)
if image is not None and img_to_img_strength > 0:
# run image2image
if not (ori_width == width and ori_height == height):
image = F.interpolate(image, (height, width), mode='bicubic')
with torch.autocast('cuda', enabled=False):
z = self.encode_first_stage(image, **kwargs)
else:
z = None
# UNet use input n_prompt
samples = self.diffusion.sample(
noise=noise,
x=z,
denoising_strength=img_to_img_strength if z is not None else 1.0,
refine_strength=refine_strength,
solver=sampler,
model=self.model,
model_kwargs=[{
'cond': context
}, {
'cond': null_context
}],
steps=sample_steps,
guide_scale=guide_scale,
guide_rescale=guide_rescale,
discretization=discretization,
show_progress=True,
seed=seed,
condition_fn=None,
clamp=None,
percentile=None,
t_max=None,
t_min=None,
discard_penultimate_step=None,
return_intermediate=None,
**kwargs)
# apply refiner
if refine_strength > 0:
assert self.refiner_model is not None
assert self.refiner_cond_model is not None
with torch.autocast('cuda', enabled=False):
before_refiner_samples = self.decode_first_stage(
samples).float()
before_refiner_samples = torch.clamp(
(before_refiner_samples + 1.0) / 2.0, min=0.0, max=1.0)
if len(pre_batch) > 0:
batch = {'prompt': prompt}
batch.update(pre_batch)
batch_uc = {'prompt': n_prompt}
batch_uc.update(pre_batch)
else:
batch, batch_uc = self.get_batch(self.input_refiner_keys,
value_dict,
num_samples=num_samples)
context = getattr(self.refiner_cond_model, 'encode')(batch)
null_context = getattr(self.refiner_cond_model, 'encode')(batch_uc)
samples = self.diffusion.sample(
noise=noise,
x=samples,
denoising_strength=img_to_img_strength
if z is not None else 1.0,
refine_strength=refine_strength,
refine_stage=True,
solver=sampler,
model=self.refiner_model,
model_kwargs=[{
'cond': context
}, {
'cond': null_context
}],
steps=sample_steps,
guide_scale=guide_scale,
guide_rescale=guide_rescale,
discretization=discretization,
show_progress=True,
seed=seed,
condition_fn=None,
clamp=None,
percentile=None,
t_max=None,
t_min=None,
discard_penultimate_step=None,
return_intermediate=None,
**kwargs)
else:
before_refiner_samples = [None for _ in prompt]
with torch.autocast('cuda', enabled=False):
x_samples = self.decode_first_stage(samples).float()
x_samples = torch.clamp((x_samples + 1.0) / 2.0, min=0.0, max=1.0)
# UNet use train n_prompt
if not self.default_n_prompt == self.train_n_prompt and run_train_n:
train_n_prompt = [self.train_n_prompt] * len(prompt)
if len(pre_batch) > 0:
pre_batch = {'prompt': prompt}
batch.update(pre_batch)
batch_uc = {'prompt': train_n_prompt}
batch_uc.update(pre_batch)
else:
value_dict['negative_prompt'] = train_n_prompt
batch, batch_uc = self.get_batch(self.input_keys,
value_dict,
num_samples=num_samples)
context = getattr(self.cond_stage_model, 'encode')(batch)
null_context = getattr(self.cond_stage_model, 'encode')(batch_uc)
tn_samples = self.diffusion.sample(
noise=noise,
x=z,
denoising_strength=img_to_img_strength
if z is not None else 1.0,
refine_strength=refine_strength,
solver=sampler,
model=self.model,
model_kwargs=[{
'cond': context
}, {
'cond': null_context
}],
steps=sample_steps,
guide_scale=guide_scale,
guide_rescale=guide_rescale,
discretization=discretization,
show_progress=we.rank == 0,
seed=seed,
condition_fn=None,
clamp=None,
percentile=None,
t_max=None,
t_min=None,
discard_penultimate_step=None,
return_intermediate=None,
**kwargs)
if refine_strength > 0:
assert self.refiner_model is not None
assert self.refiner_cond_model is not None
with torch.autocast('cuda', enabled=False):
before_refiner_t_samples = self.decode_first_stage(
samples).float()
before_refiner_t_samples = torch.clamp(
(before_refiner_t_samples + 1.0) / 2.0, min=0.0, max=1.0)
if len(pre_batch) > 0:
pre_batch = {'prompt': prompt}
batch.update(pre_batch)
batch_uc = {'prompt': train_n_prompt}
batch_uc.update(pre_batch)
else:
batch, batch_uc = self.get_batch(self.input_refiner_keys,
value_dict,
num_samples=num_samples)
context = getattr(self.refiner_cond_model, 'encode')(batch)
null_context = getattr(self.refiner_cond_model,
'encode')(batch_uc)
tn_samples = self.diffusion.sample(
noise=noise,
x=tn_samples,
denoising_strength=img_to_img_strength
if z is not None else 1.0,
refine_strength=refine_strength,
refine_stage=True,
solver=sampler,
model=self.refiner_model,
model_kwargs=[{
'cond': context
}, {
'cond': null_context
}],
steps=sample_steps,
guide_scale=guide_scale,
guide_rescale=guide_rescale,
discretization=discretization,
show_progress=True,
seed=seed,
condition_fn=None,
clamp=None,
percentile=None,
t_max=None,
t_min=None,
discard_penultimate_step=None,
return_intermediate=None,
**kwargs)
else:
before_refiner_t_samples = [None for _ in prompt]
t_x_samples = self.decode_first_stage(tn_samples).float()
t_x_samples = torch.clamp((t_x_samples + 1.0) / 2.0,
min=0.0,
max=1.0)
else:
train_n_prompt = ['' for _ in prompt]
t_x_samples = [None for _ in prompt]
before_refiner_t_samples = [None for _ in prompt]
outputs = list()
for p, np, tnp, img, r_img, t_img, r_t_img in zip(
prompt, n_prompt, train_n_prompt, x_samples,
before_refiner_samples, t_x_samples, before_refiner_t_samples):
one_tup = {
'prompt': p,
'n_prompt': np,
'image': img,
'before_refiner_image': r_img
}
if t_img is not None:
one_tup['train_n_prompt'] = tnp
one_tup['train_n_image'] = t_img
one_tup['train_n_before_refiner_image'] = r_t_img
outputs.append(one_tup)
return outputs
@staticmethod
def get_config_template():
return dict_to_yaml('MODEL',
__class__.__name__,
LatentDiffusionXL.para_dict,
set_name=True)
@@ -0,0 +1,49 @@
# -*- coding: utf-8 -*-
# Copyright (c) Alibaba, Inc. and its affiliates.
from abc import ABCMeta, abstractmethod
import torch
from scepter.modules.model.base_model import BaseModel
from scepter.modules.utils.config import dict_to_yaml
class TrainModule(BaseModel, metaclass=ABCMeta):
para_dict = {}
def __init__(self, cfg, logger=None):
super(TrainModule, self).__init__(cfg, logger=logger)
self.logger = logger
self.cfg = cfg
@abstractmethod
def forward(self, *inputs, **kwargs):
pass
@abstractmethod
def forward_train(self, *inputs, **kwargs):
pass
@abstractmethod
@torch.no_grad()
def forward_test(self, *inputs, **kwargs):
pass
@staticmethod
def get_config_template():
'''
{ "ENV" :
{ "description" : "",
"A" : {
"value": 1.0,
"description": ""
}
}
}
:return:
'''
return dict_to_yaml('networkname',
__class__.__name__,
TrainModule.para_dict,
set_name=True)