new project from v0.0.1
This commit is contained in:
@@ -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)
|
||||
@@ -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
|
||||
@@ -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)
|
||||
@@ -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)
|
||||
Reference in New Issue
Block a user