v1.0.3 update
This commit is contained in:
@@ -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
|
||||
|
||||
@@ -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)
|
||||
Reference in New Issue
Block a user