v1.0.3 update

This commit is contained in:
zeyinzi.jzyz
2024-07-18 14:12:42 +08:00
parent 7a9f90efb2
commit 01fd8335af
94 changed files with 8776 additions and 417 deletions
+2 -1
View File
@@ -4,4 +4,5 @@ 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_sce, ldm_xl
from scepter.modules.model.network.ldm import (ldm, ldm_edit, ldm_pixart,
ldm_sce, ldm_sd3, ldm_xl)
@@ -1,10 +1,11 @@
# -*- coding: utf-8 -*-
# Copyright (c) Alibaba, Inc. and its affiliates.
import re
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
@@ -87,6 +88,7 @@ class AutoencoderKL(TrainModule):
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.use_conv = self.cfg.get('USE_CONV', True)
self.construct_network()
self.init_network()
@@ -95,8 +97,12 @@ class AutoencoderKL(TrainModule):
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)
self.conv1 = torch.nn.Conv2d(
2 * z_channels, 2 *
self.embed_dim, 1) if self.use_conv else torch.nn.Identity()
self.conv2 = torch.nn.Conv2d(
self.embed_dim, z_channels,
1) if self.use_conv else torch.nn.Identity()
if self.loss_cfg is not None:
self.loss = LOSSES.build(self.loss_cfg, logger=self.logger)
@@ -107,7 +113,7 @@ class AutoencoderKL(TrainModule):
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()):
def init_from_ckpt(self, path, ignore_keys):
if path.find('.safetensors') > -1:
from safetensors import safe_open
sd = OrderedDict()
@@ -122,20 +128,17 @@ class AutoencoderKL(TrainModule):
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
if self.ignore_keys is not None:
if (isinstance(self.ignore_keys, str) and re.match(self.ignore_keys, k)) or \
(isinstance(self.ignore_keys, list) and k in self.ignore_keys):
continue
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
k = k.replace('first_stage_model.', '')
new_sd[k] = v
missing, unexpected = self.load_state_dict(new_sd, strict=False)
if we.rank == 0:
@@ -264,3 +267,14 @@ class AutoencoderKL(TrainModule):
__class__.__name__,
AutoencoderKL.para_dict,
set_name=True)
if __name__ == '__main__':
import argparse
from scepter.modules.utils.config import Config
from scepter.modules.utils.logger import get_logger
std_logger = get_logger(name='scepter')
parser = argparse.ArgumentParser(description='Argparser for Scepter:\n')
cfg = Config(load=True, parser_ins=parser)
model = AutoencoderKL(cfg, logger=std_logger)
model.load_pretrained_model(cfg.PRETRAINED_MODEL)
@@ -170,6 +170,45 @@ def adaptive_anisotropic_filter(x, g=None):
return y
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)
def get_scalings_for_boundary_condition(sigma):
sigma_data = 0.5
c_skip = (1 -
sigma**2)**0.5 * sigma_data**2 / (sigma**2 +
(1 - sigma**2) * sigma_data**2)
c_out = (sigma * sigma_data / (sigma**2 +
(1 - sigma**2) * sigma_data**2)**0.5)
return c_skip, c_out
def v_to_x0(v, t, x_t, diffusion):
sigmas = _i(diffusion.sigmas, t, v)
alphas = _i(diffusion.alphas, t, v)
return alphas * x_t - sigmas * v
class GaussianDiffusion(object):
def __init__(self, sigmas, prediction_type='eps'):
assert prediction_type in {'x0', 'eps', 'v'}
@@ -666,40 +705,212 @@ class GaussianDiffusion(object):
noise)
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)))
class GaussianDiffusionRF(object):
def __init__(self, sigmas, prediction_type='rf'):
assert prediction_type in {'rf'}
self.sigmas = sigmas
self.num_timesteps = len(sigmas)
def diffuse(self, x0, t, noise, sigma):
"""
Add Gaussian noise to signal x0 according to:
q(x_t | x_0) = N(x_t | alpha_t x_0, sigma_t^2 I).
"""
shape = (x0.size(0), ) + (1, ) * (x0.ndim - 1)
sigma = sigma.view(shape)
alpha = 1 - sigma
xt = alpha * x0 + sigma * noise
return xt
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)
def denoise(self,
xt,
t,
sigma,
model,
model_kwargs={},
guide_scale=None,
guide_rescale=None,
cat_uc=False,
**kwargs):
assert sigma is not None
shape = (xt.size(0), ) + (1, ) * (xt.ndim - 1)
sigma = sigma.view(shape)
def get_scalings_for_boundary_condition(sigma):
sigma_data = 0.5
c_skip = (1 -
sigma**2)**0.5 * sigma_data**2 / (sigma**2 +
(1 - sigma**2) * sigma_data**2)
c_out = (sigma * sigma_data / (sigma**2 +
(1 - sigma**2) * sigma_data**2)**0.5)
return c_skip, c_out
# prediction
if guide_scale is None:
if isinstance(model_kwargs, dict):
out = model(xt, t=t, **model_kwargs, **kwargs)
elif isinstance(model_kwargs, list) and len(model_kwargs) > 0:
out = model(xt, t=t, **model_kwargs[0], **kwargs)
else:
raise Exception('Error')
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 isinstance(guide_scale, float) or isinstance(guide_scale, int):
assert len(model_kwargs) == 2
if guide_scale == 1.:
out = model(xt, t=t, **model_kwargs[0], **kwargs)
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
def v_to_x0(v, t, x_t, diffusion):
sigmas = _i(diffusion.sigmas, t, v)
alphas = _i(diffusion.alphas, t, v)
return alphas * x_t - sigmas * v
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,
**kwargs)
y_out, u_out = all_out.chunk(2)
else:
y_out = model(xt, t=t, **model_kwargs[0], **kwargs)
u_out = model(xt, t=t, **model_kwargs[1], **kwargs)
out = u_out + guide_scale * (y_out - u_out)
if guide_rescale is not None and guide_rescale > 0.0:
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
x0 = xt - sigma * out
return x0
def loss(self,
x0,
t,
model,
model_kwargs={},
reduction='mean',
noise=None,
**kwargs):
sigma = t / self.num_timesteps
shape = (x0.size(0), ) + (1, ) * (x0.ndim - 1)
sigma = sigma.view(shape)
if noise is None:
noise = torch.randn_like(x0)
xt = self.diffuse(x0, t, noise, sigma=sigma)
out = model(xt, t=t, **model_kwargs, **kwargs)
loss = ((xt - sigma * out) - x0)**2
# loss = (out - (x0 - noise)) ** 2
if reduction == 'mean':
loss = loss.flatten(1).mean(dim=1)
return loss
@torch.no_grad()
def sample(self,
noise,
model,
model_kwargs={},
guide_scale=None,
guide_rescale=None,
solver='euler',
steps=20,
shift=3,
discretization=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 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,
'onestep': sample_onestep,
'multistep': stochastic_iterative_sampler,
'multistep2': stochastic_iterative_sampler2,
'multistep3': stochastic_iterative_sampler3,
'dpmpp_2m_sde_lcm': sample_dpmpp_2m_sde_lcm,
}[solver]
seed = seed if seed >= 0 else random.randint(0, 2**31)
intermediates = []
def model_fn(xt, sigma):
# denoising
sigma = sigma.repeat(len(xt)).to(xt.device)
t = self._sigma_to_t(sigma).round().long()
x0 = self.denoise(xt,
t,
sigma,
model,
model_kwargs,
guide_scale,
guide_rescale,
cat_uc=cat_uc,
**kwargs)
# 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
device = self.sigmas.device
sigma_max = self.sigmas[0]
sigma_min = self.sigmas[-1]
t_max = sigma_max * self.num_timesteps
t_min = sigma_min * self.num_timesteps
steps = torch.linspace(t_max, t_min, steps).to(device)
sigmas = steps / self.num_timesteps
sigmas = shift * sigmas / (1 + (shift - 1) * sigmas)
sigmas = sigmas.to(torch.float32).to(device)
sigmas = torch.cat([sigmas, sigmas.new_zeros([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):
return sigma * self.num_timesteps
@@ -13,6 +13,7 @@ where alpha_t^2 = 1 - sigma_t^2.
"""
import math
import numpy as np
import torch
__all__ = [
@@ -154,6 +155,15 @@ def logsnr_cosine_interp_schedule(n,
_logsnr_cosine_interp(n, logsnr_min, logsnr_max, scale_min, scale_max))
def shifted_schedule(n, shift=3):
timesteps = np.linspace(1, n, n, dtype=np.float32)[::-1].copy()
timesteps = torch.from_numpy(timesteps).to(dtype=torch.float32)
sigmas = timesteps / n
sigmas = shift * sigmas / (1 + (shift - 1) * sigmas)
return sigmas
def noise_schedule(schedule='logsnr_cosine_interp',
n=1000,
zero_terminal_snr=False,
@@ -171,7 +181,8 @@ def noise_schedule(schedule='logsnr_cosine_interp',
'vp': vp_schedule,
'logsnr_cosine': logsnr_cosine_schedule,
'logsnr_cosine_shifted': logsnr_cosine_shifted_schedule,
'logsnr_cosine_interp': logsnr_cosine_interp_schedule
'logsnr_cosine_interp': logsnr_cosine_interp_schedule,
'shifted': shifted_schedule,
}[schedule](n, **kwargs)
# post-processing
@@ -12,9 +12,8 @@ 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.
"""
from tqdm.auto import trange
import torch
from tqdm.auto import trange
__all__ = [
'sample_euler', 'sample_euler_ancestral', 'sample_heun', 'sample_dpm_2',
@@ -77,8 +76,9 @@ def sample_euler(noise,
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)
# _, c_in = get_scalings(sigma_hat)
# denoised = model(x * c_in, sigma_hat)
denoised = model(x, sigmas[i])
d = (x - denoised) / sigma_hat
dt = sigmas[i + 1] - sigma_hat
x = x + d * dt
@@ -2,7 +2,9 @@
# Copyright (c) Alibaba, Inc. and its affiliates.
from scepter.modules.model.network.ldm.ldm import LatentDiffusion
from scepter.modules.model.network.ldm.ldm_edit import LatentDiffusionEdit
from scepter.modules.model.network.ldm.ldm_pixart import LatentDiffusionPixart
from scepter.modules.model.network.ldm.ldm_sce import (
LatentDiffusionSCEControl, LatentDiffusionSCETuning,
LatentDiffusionXLSCEControl, LatentDiffusionXLSCETuning)
from scepter.modules.model.network.ldm.ldm_sd3 import LatentDiffusionSD3
from scepter.modules.model.network.ldm.ldm_xl import LatentDiffusionXL
+22 -15
View File
@@ -1,10 +1,12 @@
# -*- coding: utf-8 -*-
# Copyright (c) Alibaba, Inc. and its affiliates.
import copy
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
@@ -91,8 +93,8 @@ class LatentDiffusion(TrainModule):
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"'
'eps', 'x0', 'v', 'rf'
], 'currently only supporting "eps" and "x0" and "v" and "rf"'
self.num_timesteps = self.cfg.get('TIMESTEPS', 1000)
self.schedule_args = {
@@ -137,15 +139,18 @@ class LatentDiffusion(TrainModule):
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.use_ema = self.cfg.get('USE_EMA', False)
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)
if self.use_ema:
if self.model_ema_config:
self.model_ema = BACKBONES.build(self.model_ema_config,
logger=self.logger)
else:
self.model_ema = copy.deepcopy(self.model)
self.model_ema = self.model_ema.eval()
for param in self.model_ema.parameters():
param.requires_grad = False
@@ -154,13 +159,15 @@ class LatentDiffusion(TrainModule):
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.first_stage_config:
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
else:
self.first_stage_model = None
if self.tokenizer_config is not None:
self.cond_stage_config.KWARGS = {
'vocab_size': self.tokenizer.vocab_size
@@ -269,8 +276,8 @@ class LatentDiffusion(TrainModule):
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,
def noise_sample(self, batch_size, h, w, g, c=4):
noise = torch.empty(batch_size, c, h, w,
device=we.device_id).normal_(generator=g)
return noise
@@ -0,0 +1,272 @@
# -*- coding: utf-8 -*-
# Copyright (c) Alibaba, Inc. and its affiliates.
import copy
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.ldm import LatentDiffusion
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 LatentDiffusionPixart(LatentDiffusion):
para_dict = LatentDiffusion.para_dict
para_dict['DECODER_BIAS'] = {'value': 0, 'description': ''}
def __init__(self, cfg, logger=None):
super().__init__(cfg, logger=logger)
self.decoder_bias = cfg.get('DECODER_BIAS', 0.5)
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:
self.model_ema = copy.deepcopy(self.model).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)
if self.first_stage_config:
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
else:
self.first_stage_model = None
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 forward_train(self,
image=None,
noise=None,
prompt=None,
label=None,
**kwargs):
n, c, h, w = image.shape
x_start = self.encode_first_stage(image, **kwargs)
t = torch.randint(0, self.num_timesteps, (n, ),
device=x_start.device).long()
ar = torch.tensor([[h / w]], device=we.device_id).repeat(n, 1)
hw = torch.tensor([[h, w]], dtype=torch.float,
device=we.device_id).repeat(n, 1)
context = {}
cont_mask = None
if prompt and self.cond_stage_model:
with torch.autocast(device_type='cuda',
enabled=True,
dtype=torch.bfloat16):
cont, cont_mask = getattr(self.cond_stage_model,
'encode')(prompt, return_mask=True)
context['crossattn'] = cont.float()
else:
assert label is not None
context['label'] = label
if 'hint' in kwargs and kwargs['hint'] is not None:
hint = kwargs.pop('hint')
if isinstance(context, dict):
context['hint'] = hint
else:
context = {'crossattn': context, 'hint': hint}
else:
hint = None
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,
'mask': cont_mask,
'data_info': {
'img_hw': hw,
'aspect_ratio': ar
}
},
noise=noise,
**kwargs)
loss = loss * weights
loss = loss.mean()
ret = {'loss': loss, 'probe_data': {'prompt': prompt}}
return ret
@torch.no_grad()
def forward_test(self,
prompt=None,
label=None,
sampler='ddim',
sample_steps=20,
seed=2023,
guide_scale=4.5,
guide_rescale=0.5,
discretization='trailing',
**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 = label.shape[0] if label is not None else len(prompt)
context = {}
null_context = {}
cont_mask = None
if prompt and self.cond_stage_model:
with torch.autocast(device_type='cuda',
enabled=True,
dtype=torch.bfloat16):
cont, cont_mask = getattr(self.cond_stage_model,
'encode')(prompt, return_mask=True)
context['crossattn'] = cont.float()
null_context['crossattn'] = self.model.y_embedder.y_embedding[
None].repeat(len(prompt), 1, 1)
else:
assert label is not None
context['label'] = label
null_context['label'] = torch.tensor(
[self.model.num_classes]).repeat(num_samples).to(we.device_id)
if 'hint' in kwargs and kwargs['hint'] is not None:
hint = kwargs.pop('hint')
if isinstance(context, dict):
context['hint'] = hint
if isinstance(null_context, dict):
null_context['hint'] = hint
else:
hint = None
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 isinstance(image_size, numbers.Number):
image_size = [image_size, image_size]
if image_size is None:
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,
'mask': cont_mask,
'data_info': {
'img_hw':
torch.tensor([image_size],
dtype=torch.float,
device=we.device_id).repeat(num_samples, 1),
'aspect_ratio':
torch.tensor([[1.]], device=we.device_id).repeat(
num_samples, 1)
}
}, {
'cond': null_context,
'mask': cont_mask,
'data_info': {
'img_hw':
torch.tensor([image_size],
dtype=torch.float,
device=we.device_id).repeat(num_samples, 1),
'aspect_ratio':
torch.tensor([[1.]], device=we.device_id).repeat(
num_samples, 1)
}
}] if guide_scale is not None and guide_scale > 0 else {
'cond': context,
'mask': cont_mask,
'data_info': {
'img_hw':
torch.tensor([image_size],
dtype=torch.float,
device=we.device_id).repeat(num_samples, 1),
'aspect_ratio':
torch.tensor([[1.]], device=we.device_id).repeat(
num_samples, 1)
}
},
cat_uc=False,
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 + self.decoder_bias / 255,
min=0.0,
max=1.0)
outputs = list()
prompt = label.detach().cpu().numpy().tolist(
) if prompt is None else prompt
for i, (p, img) in enumerate(zip(prompt, x_samples)):
one_tup = {'prompt': str(p), 'n_prompt': '', 'image': img}
if hint is not None:
one_tup.update({'hint': hint[i]})
outputs.append(one_tup)
return outputs
@staticmethod
def get_config_template():
return dict_to_yaml('MODEL',
__class__.__name__,
LatentDiffusionPixart.para_dict,
set_name=True)
@@ -0,0 +1,236 @@
# -*- coding: utf-8 -*-
# Copyright (c) Alibaba, Inc. and its affiliates.
import numbers
import random
import torch
import torch.nn.functional as F
from scepter.modules.model.network.diffusion.diffusion import \
GaussianDiffusionRF
from scepter.modules.model.network.diffusion.schedules import noise_schedule
from scepter.modules.model.network.ldm import LatentDiffusion
from scepter.modules.model.registry import MODELS
from scepter.modules.utils.config import dict_to_yaml
from scepter.modules.utils.distribute import we
@MODELS.register_class()
class LatentDiffusionSD3(LatentDiffusion):
para_dict = LatentDiffusion.para_dict
def __init__(self, cfg, logger=None):
super().__init__(cfg, logger=logger)
self.shift_factor = cfg.get('SHIFT_FACTOR', 0)
self.t_weight_type = cfg.get('T_WEIGHT', 'logit_normal')
self.logit_mean = cfg.get('LOGIT_MEAN', 0.0)
self.logit_std = cfg.get('LOGIT_STD', 1.0)
def init_params(self):
self.parameterization = self.cfg.get('PARAMETERIZATION', 'rf')
assert self.parameterization in [
'eps', 'x0', 'v', 'rf'
], 'currently only supporting "eps" and "x0" and "v" and "rf"'
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 = GaussianDiffusionRF(
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', False)
self.model_ema_config = self.cfg.get('DIFFUSION_MODEL_EMA', None)
def noise_sample(self, batch_size, h, w, g, c=4):
noise = torch.empty(batch_size, c, h, w,
device=we.device_id).normal_(generator=g)
return noise
def forward_train(self, image=None, noise=None, prompt=None, **kwargs):
n, c, h, w = image.shape
x_start = self.encode_first_stage(image, **kwargs)
if self.t_weight_type == 'uniform':
t = torch.randint(0,
self.num_timesteps, (n, ),
device=x_start.device).long()
elif self.t_weight_type == 'logit_normal':
density = F.sigmoid(
torch.normal(mean=self.logit_mean,
std=self.logit_std,
size=(n, ),
device=x_start.device))
t = (density * (self.num_timesteps - 1)).round().long()
sigma = (t + 1) / self.num_timesteps
shift = self.schedule_args['shift']
if shift > 1.:
sigma = shift * sigma / (1 + (shift - 1) * sigma)
t = sigma * self.num_timesteps
context = {}
if prompt and self.cond_stage_model:
ctx, pooled = getattr(self.cond_stage_model, 'encode')(prompt)
context['crossattn'] = ctx.float()
context['y'] = pooled
else:
assert False
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,
**kwargs)
loss = loss * weights
loss = loss.mean()
ret = {'loss': loss, 'probe_data': {'prompt': prompt}}
return ret
@torch.no_grad()
def forward_test(self,
prompt=None,
sampler='ddim',
sample_steps=20,
seed=2023,
guide_scale=4.5,
guide_rescale=0.0,
**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)
context = {}
null_context = {}
if prompt and self.cond_stage_model:
ctx, pooled = getattr(self.cond_stage_model, 'encode')(prompt)
null_ctx, null_pooled = getattr(self.cond_stage_model,
'encode')([''] * len(prompt))
context['crossattn'] = ctx.float()
context['y'] = pooled.float()
null_context['crossattn'] = null_ctx.float()
null_context['y'] = null_pooled.float()
else:
assert False
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 isinstance(image_size, numbers.Number):
image_size = [image_size, image_size]
if image_size is None:
image_size = [1024, 1024]
height, width = image_size
noise = self.noise_sample(num_samples,
height // self.size_factor,
width // self.size_factor,
g,
c=16)
# UNet use input n_prompt
samples = self.diffusion.sample(
solver=sampler,
noise=noise,
model=self.model,
model_kwargs=[{
'cond': context
}, {
'cond': null_context
}] if guide_scale is not None and guide_scale > 0 else {
'cond': context,
},
cat_uc=False,
steps=sample_steps,
guide_scale=guide_scale,
guide_rescale=guide_rescale,
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)
outputs = list()
for i, (p, img) in enumerate(zip(prompt, x_samples)):
one_tup = {'prompt': str(p), 'n_prompt': '', 'image': img}
outputs.append(one_tup)
return outputs
@staticmethod
def get_config_template():
return dict_to_yaml('MODEL',
__class__.__name__,
LatentDiffusionSD3.para_dict,
set_name=True)
@torch.no_grad()
def encode_first_stage(self, x, **kwargs):
z = self.first_stage_model.encode(x)
return self.scale_factor * (z - self.shift_factor)
@torch.no_grad()
def decode_first_stage(self, z):
z = 1. / self.scale_factor * z + self.shift_factor
return self.first_stage_model.decode(z)