update v0.0.4

This commit is contained in:
LouieStark
2024-03-31 13:08:41 +08:00
parent 35aada8ce8
commit bf53829530
106 changed files with 6927 additions and 889 deletions
@@ -1,4 +1,5 @@
# -*- coding: utf-8 -*-
# Copyright (c) Alibaba, Inc. and its affiliates.
from scepter.modules.model.backbone.autoencoder.ae_module import (Decoder,
Encoder)
Encoder,
RDecoder)
@@ -3,6 +3,7 @@
import torch
import torch.nn as nn
from einops import repeat
from torch.utils.checkpoint import checkpoint
from scepter.modules.model.backbone.autoencoder.ae_utils import (
@@ -244,6 +245,7 @@ class Decoder(BaseModel):
# compute in_ch_mult, block_in and curr_res at lowest res
block_in = self.ch * self.ch_mult[self.num_resolutions - 1]
self.block_in = block_in
curr_res = 1
# z to block_in
self.conv_in = torch.nn.Conv2d(self.z_channels,
@@ -340,3 +342,48 @@ class Decoder(BaseModel):
__class__.__name__,
Decoder.para_dict,
set_name=True)
@BACKBONES.register_class()
class RDecoder(Decoder):
def construct_model(self):
super().construct_model()
self.resize_level = nn.Sequential(
nn.Linear(self.block_in, self.block_in),
nn.SiLU(),
nn.Linear(self.block_in, self.block_in),
)
def forward(self, z, rembed=None):
# timestep embedding
temb = None
h = self.conv_in(z)
bs, channel, hdim, wdim = h.size()
if rembed is not None:
rembed = self.resize_level(rembed)
rembed = repeat(rembed, 'b e-> b e hd wd', hd=hdim, wd=wdim)
h = h + rembed
# middle
if not self.use_checkpoint:
h = self.mid_upsclae_transform(h, temb)
else:
h = checkpoint(self.mid_upsclae_transform, h, temb)
# end
if self.give_pre_end:
return h
h = self.norm_out(h)
h = nonlinearity(h)
h = self.conv_out(h)
if self.tanh_out:
h = torch.tanh(h)
return h
@staticmethod
def get_config_template():
return dict_to_yaml('BACKBONE',
__class__.__name__,
Decoder.para_dict,
set_name=True)
@@ -1,3 +1,5 @@
# -*- coding: utf-8 -*-
# Copyright (c) Alibaba, Inc. and its affiliates.
from scepter.modules.model.backbone.unet.unet_module import DiffusionUNet
from scepter.modules.model.backbone.unet.unet_module import (DiffusionUNet,
DiffusionUNetXL,
LargenUNetXL)
@@ -8,8 +8,9 @@ import torch
import torch.nn as nn
from scepter.modules.model.backbone.unet.unet_utils import (
Downsample, ResBlock, SpatialTransformer, Timestep,
TimestepEmbedSequential, Upsample, conv_nd, linear, normalization,
BasicTransformerBlock, Downsample, ResBlock, SpatialTransformer,
SpatialTransformerV2, Timestep, TimestepEmbedSequential,
TransformerBlockV2, Upsample, conv_nd, linear, normalization,
timestep_embedding, zero_module)
from scepter.modules.model.base_model import BaseModel
from scepter.modules.model.registry import BACKBONES
@@ -951,3 +952,443 @@ class DiffusionUNetXL(DiffusionUNet):
__class__.__name__,
DiffusionUNetXL.para_dict,
set_name=True)
@BACKBONES.register_class()
class LargenUNetXL(DiffusionUNetXL):
para_dict = {
'TRANSFORMER_BLOCK_TYPE': {
'value': 'att_v1'
},
'IMAGE_SCALE': {
'value': 0.0,
},
}
para_dict.update(DiffusionUNetXL.para_dict)
def __init__(self, cfg, logger):
super().__init__(cfg, logger=logger)
self.init_params(cfg)
self.construct_network()
def init_params(self, cfg):
super().init_params(cfg)
self.transformer_block_type = cfg.get('TRANSFORMER_BLOCK_TYPE',
'att_v1')
TRANSFORMER_BLOCKS = {
'att_v1': BasicTransformerBlock,
'att_v2': TransformerBlockV2,
}
assert self.transformer_block_type in list(TRANSFORMER_BLOCKS.keys())
self.transformer_block = TRANSFORMER_BLOCKS[
self.transformer_block_type]
self.image_scale = cfg.get('IMAGE_SCALE', 0.0)
self.use_refine = cfg.get('USE_REFINE', False)
def construct_network(self):
in_channels = self.in_channels
model_channels = self.model_channels
out_channels = self.out_channels
attention_resolutions = self.attention_resolutions
channel_mult = self.channel_mult
num_classes = self.num_classes
num_heads = self.num_heads
num_head_channels = self.num_head_channels
dims = self.dims
dropout = self.dropout
use_checkpoint = self.use_checkpoint
use_scale_shift_norm = self.use_scale_shift_norm
disable_self_attentions = self.disable_self_attentions
disable_middle_self_attn = self.disable_middle_self_attn
transformer_depth = self.transformer_depth
transformer_depth_middle = self.transformer_depth_middle
context_dim = self.context_dim
use_linear_in_transformer = self.use_linear_in_transformer
resblock_updown = self.resblock_updown
conv_resample = self.conv_resample
adm_in_channels = self.adm_in_channels
transformer_block = self.transformer_block
time_embed_dim = model_channels * 4
self.time_embed = nn.Sequential(
linear(model_channels, time_embed_dim),
nn.SiLU(),
linear(time_embed_dim, time_embed_dim),
)
if self.num_classes is not None:
if isinstance(self.num_classes, int):
self.label_emb = nn.Embedding(num_classes, time_embed_dim)
elif self.num_classes == 'continuous':
print('setting up linear c_adm embedding layer')
self.label_emb = nn.Linear(1, time_embed_dim)
elif self.num_classes == 'timestep':
self.label_emb = nn.Sequential(
Timestep(model_channels),
nn.Sequential(
linear(model_channels, time_embed_dim),
nn.SiLU(),
linear(time_embed_dim, time_embed_dim),
),
)
elif self.num_classes == 'sequential':
assert adm_in_channels is not None
self.label_emb = nn.Sequential(
nn.Sequential(
linear(adm_in_channels, time_embed_dim),
nn.SiLU(),
linear(time_embed_dim, time_embed_dim),
))
else:
raise ValueError()
self.input_blocks = nn.ModuleList([
TimestepEmbedSequential(
conv_nd(dims, in_channels, model_channels, 3, padding=1))
])
self._feature_size = model_channels
input_block_chans = [model_channels]
input_down_flag = [False]
ch = model_channels
ds = 1
for level, mult in enumerate(channel_mult):
for nr in range(self.num_res_blocks[level]):
layers = [
ResBlock(
ch,
time_embed_dim,
dropout,
out_channels=mult * model_channels,
dims=dims,
use_checkpoint=use_checkpoint,
use_scale_shift_norm=use_scale_shift_norm,
)
]
ch = mult * model_channels
if ds in attention_resolutions:
if num_head_channels == -1:
dim_head = ch // num_heads
else:
num_heads = ch // num_head_channels
dim_head = num_head_channels
disabled_sa = disable_self_attentions[level] if exists(
disable_self_attentions) else False
layers.append(
SpatialTransformerV2(
ch,
num_heads,
dim_head,
transformer_block=transformer_block,
depth=transformer_depth[level],
context_dim=context_dim,
disable_self_attn=disabled_sa,
use_linear=use_linear_in_transformer,
use_checkpoint=use_checkpoint))
self.input_blocks.append(TimestepEmbedSequential(*layers))
self._feature_size += ch
input_block_chans.append(ch)
input_down_flag.append(False)
if level != len(channel_mult) - 1:
out_ch = ch
self.input_blocks.append(
TimestepEmbedSequential(
ResBlock(
ch,
time_embed_dim,
dropout,
out_channels=out_ch,
dims=dims,
use_checkpoint=use_checkpoint,
use_scale_shift_norm=use_scale_shift_norm,
down=True,
) if resblock_updown else Downsample(
ch, conv_resample, dims=dims, out_channels=out_ch))
)
ch = out_ch
input_block_chans.append(ch)
input_down_flag.append(True)
ds *= 2
self._feature_size += ch
self._input_block_chans = copy.deepcopy(input_block_chans)
self._input_down_flag = input_down_flag
if num_head_channels == -1:
dim_head = ch // num_heads
else:
num_heads = ch // num_head_channels
dim_head = num_head_channels
self.middle_block = TimestepEmbedSequential(
ResBlock(
ch,
time_embed_dim,
dropout,
dims=dims,
use_checkpoint=use_checkpoint,
use_scale_shift_norm=use_scale_shift_norm,
),
SpatialTransformerV2(ch,
num_heads,
dim_head,
transformer_block=transformer_block,
depth=transformer_depth_middle,
context_dim=context_dim,
disable_self_attn=disable_middle_self_attn,
use_linear=use_linear_in_transformer,
use_checkpoint=use_checkpoint),
ResBlock(
ch,
time_embed_dim,
dropout,
dims=dims,
use_checkpoint=use_checkpoint,
use_scale_shift_norm=use_scale_shift_norm,
),
)
self._feature_size += ch
self._middle_block_chans = [ch]
self._output_block_chans = []
self.output_blocks = nn.ModuleList([])
for level, mult in list(enumerate(channel_mult))[::-1]:
for i in range(self.num_res_blocks[level] + 1):
ich = input_block_chans.pop()
layers = [
ResBlock(
ch + ich,
time_embed_dim,
dropout,
out_channels=model_channels * mult,
dims=dims,
use_checkpoint=use_checkpoint,
use_scale_shift_norm=use_scale_shift_norm,
)
]
ch = model_channels * mult
if ds in attention_resolutions:
if num_head_channels == -1:
dim_head = ch // num_heads
else:
num_heads = ch // num_head_channels
dim_head = num_head_channels
disabled_sa = disable_self_attentions[level] if exists(
disable_self_attentions) else False
layers.append(
SpatialTransformerV2(
ch,
num_heads,
dim_head,
transformer_block=transformer_block,
depth=transformer_depth[level],
context_dim=context_dim,
disable_self_attn=disabled_sa,
use_linear=use_linear_in_transformer,
use_checkpoint=use_checkpoint))
if level and i == self.num_res_blocks[level]:
out_ch = ch
layers.append(
ResBlock(
ch,
time_embed_dim,
dropout,
out_channels=out_ch,
dims=dims,
use_checkpoint=use_checkpoint,
use_scale_shift_norm=use_scale_shift_norm,
up=True,
) if resblock_updown else Upsample(
ch, conv_resample, dims=dims, out_channels=out_ch))
ds //= 2
self.output_blocks.append(TimestepEmbedSequential(*layers))
self._feature_size += ch
self._output_block_chans.append(ch)
self.out = nn.Sequential(
normalization(ch),
nn.SiLU(),
zero_module(
conv_nd(dims, model_channels, out_channels, 3, padding=1)),
)
if self.use_refine:
self.ref_time_embed = copy.deepcopy(self.time_embed)
self.ref_label_emb = copy.deepcopy(self.label_emb)
self.ref_input_blocks = copy.deepcopy(self.input_blocks)
self.ref_input_blocks[0] = TimestepEmbedSequential(
conv_nd(dims, 4, model_channels, 3, padding=1))
self.ref_middle_block = copy.deepcopy(self.middle_block)
self.ref_output_blocks = copy.deepcopy(self.output_blocks)
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 == 'input_blocks.0.0.weight':
if we.rank == 0:
self.logger.info(
'Partial initial key {} from state_dict.'.format(
k))
new_v = torch.empty(320, self.in_channels, 3, 3)
nn.init.zeros_(new_v)
new_v[:, :v.shape[1]] = v
new_sd[k] = new_v
if self.use_refine:
new_sd['ref_' + k] = v
else:
new_sd[k] = v
if self.use_refine:
new_sd['ref_' + 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 forward(self, x, t=None, cond=dict(), **kwargs):
t_emb = timestep_embedding(t,
self.model_channels,
repeat_only=False,
legacy=True)
emb = self.time_embed(t_emb)
if isinstance(cond, dict):
if 'y' in cond:
assert self.num_classes is not None
emb = emb + self.label_emb(cond['y'])
if self.use_refine:
ref_emb = self.ref_time_embed(t_emb)
assert 'null_y' in cond
cond_y = cond['y'].clone()
cond_y[:, :cond['null_y'].shape[1]] = cond['null_y']
ref_emb = ref_emb + self.ref_label_emb(cond_y)
if 'concat' in cond:
c = cond['concat']
x = torch.cat([x, c], dim=1)
context = cond.get('crossattn', None)
img_context = cond.get('img_crossattn', None)
task = cond['task']
image_scale = cond.get('image_scale', self.image_scale)
if 'Subject' in task and img_context is not None:
ip_enc_scale = image_scale
ip_dec_scale = image_scale
num_img_tokens = img_context.shape[1]
context = torch.cat([context, img_context], dim=1)
else:
ip_enc_scale = None
ip_dec_scale = None
num_img_tokens = None
ref = cond.get('ref_xt', None)
ref_context = cond.get('ref_crossattn', None)
else:
raise TypeError
hs = []
refs = []
h = x
if self.use_refine:
assert ref is not None and ref_context is not None
for i, (ref_module, module) in enumerate(
zip(self.ref_input_blocks, self.input_blocks)):
ref = ref_module(ref, ref_emb, ref_context, caching=None)
h = module(h,
emb,
context,
caching=None,
scale=ip_enc_scale,
num_img_token=num_img_tokens)
refs.append(ref)
hs.append(h)
ref = self.ref_middle_block(ref,
ref_emb,
ref_context,
caching=None)
h = self.middle_block(h,
emb,
context,
caching=None,
scale=ip_enc_scale,
num_img_token=num_img_tokens)
for i, (ref_module, module) in enumerate(
zip(self.ref_output_blocks, self.output_blocks)):
cache = []
ref = torch.cat([ref, refs.pop()], dim=1)
ref = ref_module(ref,
ref_emb,
ref_context,
caching='write',
cache=cache)
h = torch.cat([h, hs.pop()], dim=1)
h = module(h,
emb,
context,
caching='read',
cache=cache,
scale=ip_dec_scale,
num_img_token=num_img_tokens)
else:
for module in self.input_blocks:
h = module(h,
emb,
context,
caching=None,
scale=ip_enc_scale,
num_img_token=num_img_tokens)
hs.append(h)
h = self.middle_block(h,
emb,
context,
caching=None,
scale=ip_enc_scale,
num_img_token=num_img_tokens)
for module in self.output_blocks:
h = torch.cat([h, hs.pop()], dim=1)
h = module(h,
emb,
context,
caching=None,
scale=ip_dec_scale,
num_img_token=num_img_tokens)
out = self.out(h)
return out
@staticmethod
def get_config_template():
return dict_to_yaml('BACKBONE',
__class__.__name__,
LargenUNetXL.para_dict,
set_name=True)
@@ -10,6 +10,7 @@ import numpy as np
import torch
import torch.nn as nn
import torch.nn.functional as F
import torchvision.transforms.functional as TF
from einops import rearrange, repeat
from packaging import version
@@ -171,12 +172,14 @@ class TimestepEmbedSequential(nn.Sequential, TimestepBlock):
A sequential module that passes timestep embeddings to the children that
support it as an extra input.
"""
def forward(self, x, emb, context=None, target_size=None):
def forward(self, x, emb, context=None, target_size=None, **kwargs):
for layer in self:
if isinstance(layer, TimestepBlock):
x = layer(x, emb)
elif isinstance(layer, SpatialTransformer):
x = layer(x, context)
elif isinstance(layer, SpatialTransformerV2):
x = layer(x, context, **kwargs)
elif isinstance(layer, Upsample):
x = layer(x, target_size)
else:
@@ -864,6 +867,92 @@ class MemoryEfficientCrossAttention(nn.Module):
return self.to_out(out)
class XFormersMHA_IP(nn.Module):
def __init__(self,
query_dim,
context_dim=None,
heads=8,
dim_head=64,
dropout=0.0):
super().__init__()
inner_dim = dim_head * heads
context_dim = default(context_dim, query_dim)
self.heads = heads
self.dim_head = dim_head
self.to_q = nn.Linear(query_dim, inner_dim, bias=False)
self.to_k = nn.Linear(context_dim, inner_dim, bias=False)
self.to_v = nn.Linear(context_dim, inner_dim, bias=False)
self.to_k_ip = nn.Linear(context_dim, inner_dim, bias=False)
self.to_v_ip = nn.Linear(context_dim, inner_dim, bias=False)
self.to_out = nn.Sequential(nn.Linear(inner_dim, query_dim),
nn.Dropout(dropout))
self.attention_op = None
def forward(self,
x,
context=None,
mask=None,
scale=None,
num_img_token=None):
q = self.to_q(x)
context = default(context, x)
if scale is not None and num_img_token is not None:
eos = context.shape[1] - num_img_token
txt_context = context[:, :eos, :]
img_context = context[:, eos:, :]
k = self.to_k(txt_context)
v = self.to_v(txt_context)
k_i = self.to_k_ip(img_context)
v_i = self.to_v_ip(img_context)
b, _, _ = q.shape
q, k, v, k_i, v_i = map(
lambda t: t.unsqueeze(3).reshape(b, t.shape[
1], self.heads, self.dim_head).permute(0, 2, 1, 3).reshape(
b * self.heads, t.shape[1], self.dim_head).contiguous(
),
(q, k, v, k_i, v_i),
)
# actually compute the attention, what we cannot get enough of
txt_out = xformers.ops.memory_efficient_attention(
q, k, v, attn_bias=None, op=self.attention_op)
img_out = xformers.ops.memory_efficient_attention(
q, k_i, v_i, attn_bias=None, op=self.attention_op)
out = txt_out + scale * img_out
else:
k = self.to_k(context)
v = self.to_v(context)
b, _, _ = q.shape
q, k, v = map(
lambda t: t.unsqueeze(3).reshape(b, t.shape[
1], self.heads, self.dim_head).permute(0, 2, 1, 3).reshape(
b * self.heads, t.shape[1], self.dim_head).contiguous(
),
(q, k, v),
)
out = xformers.ops.memory_efficient_attention(q,
k,
v,
attn_bias=None,
op=self.attention_op)
# TODO: Use this directly in the attention operation, as a bias
if exists(mask):
raise NotImplementedError
out = (out.unsqueeze(0).reshape(
b, self.heads, out.shape[1],
self.dim_head).permute(0, 2, 1,
3).reshape(b, out.shape[1],
self.heads * self.dim_head))
return self.to_out(out)
class BasicTransformerBlock(nn.Module):
def __init__(self,
dim,
@@ -908,6 +997,65 @@ class BasicTransformerBlock(nn.Module):
return x
class TransformerBlockV2(nn.Module):
def __init__(self,
query_dim,
n_heads,
d_head,
dropout=0.,
context_dim=None,
gated_ff=True,
use_checkpoint=False,
disable_self_attn=False):
super().__init__()
self.disable_self_attn = disable_self_attn
self.attn1 = MemoryEfficientCrossAttention(query_dim=query_dim,
heads=n_heads,
dim_head=d_head,
dropout=dropout,
context_dim=None)
self.ff = FeedForward(query_dim, dropout=dropout, glu=gated_ff)
self.attn2 = XFormersMHA_IP(query_dim=query_dim,
heads=n_heads,
dim_head=d_head,
context_dim=context_dim)
self.norm1 = nn.LayerNorm(query_dim)
self.norm2 = nn.LayerNorm(query_dim)
self.norm3 = nn.LayerNorm(query_dim)
self.use_checkpoint = use_checkpoint
def forward(self,
x,
context,
caching=None,
cache=None,
scale=None,
num_img_token=None,
**kwargs):
y = self.norm1(x)
if caching == 'write':
assert isinstance(cache, list)
cache.append(y)
x = self.attn1(y, context=None) + x
elif caching == 'read':
assert isinstance(cache, list) and len(cache) > 0
c = cache.pop(0)
self_ctx = torch.cat([y, c], dim=1)
x = self.attn1(y, context=self_ctx) + x
elif caching is None:
x = self.attn1(y, context=None) + x
else:
assert False
x = self.attn2(self.norm2(x),
context=context,
scale=scale,
num_img_token=num_img_token) + x
x = self.ff(self.norm3(x)) + x
return x
class SpatialTransformer(nn.Module):
"""
Transformer block for image-like data.
@@ -1003,3 +1151,108 @@ class SpatialTransformer(nn.Module):
if not self.use_linear:
x = self.proj_out(x)
return x + x_in
class SpatialTransformerV2(nn.Module):
"""
Transformer block for image-like data.
First, project the input (aka embedding)
and reshape to b, t, d.
Then apply standard transformer action.
Finally, reshape to image
NEW: use_linear for more efficiency instead of the 1x1 convs
"""
def __init__(self,
in_channels,
n_heads,
d_head,
transformer_block,
depth=1,
dropout=0.,
context_dim=None,
disable_self_attn=False,
use_linear=False,
use_checkpoint=True):
super().__init__()
if exists(context_dim) and not isinstance(context_dim, list):
context_dim = [context_dim]
if exists(context_dim) and not isinstance(context_dim, (list)):
context_dim = [context_dim]
if exists(context_dim) and isinstance(context_dim, list):
if depth != len(context_dim):
print(
f'WARNING: {self.__class__.__name__}: Found context dims {context_dim} of'
f" depth {len(context_dim)}, which does not match the specified 'depth' of"
f' {depth}. Setting context_dim to {depth * [context_dim[0]]} now.'
)
# depth does not match context dims.
assert all(
map(lambda x: x == context_dim[0], context_dim)
), 'need homogenous context_dim to match depth automatically'
context_dim = depth * [context_dim[0]]
elif context_dim is None:
context_dim = [None] * depth
self.in_channels = in_channels
inner_dim = n_heads * d_head
self.norm = normalization(in_channels)
if not use_linear:
self.proj_in = nn.Conv2d(in_channels,
inner_dim,
kernel_size=1,
stride=1,
padding=0)
else:
self.proj_in = nn.Linear(in_channels, inner_dim)
self.transformer_blocks = nn.ModuleList([
transformer_block(inner_dim,
n_heads,
d_head,
dropout=dropout,
context_dim=context_dim[d],
disable_self_attn=disable_self_attn,
use_checkpoint=use_checkpoint)
for d in range(depth)
])
if not use_linear:
self.proj_out = zero_module(
nn.Conv2d(inner_dim,
in_channels,
kernel_size=1,
stride=1,
padding=0))
else:
self.proj_out = zero_module(nn.Linear(in_channels, inner_dim))
self.use_linear = use_linear
def forward(self, x, context=None, **kwargs):
# note: if no context is given, cross-attention defaults to self-attention
if not isinstance(context, list):
context = [context]
b, c, h, w = x.shape
ref_mask = kwargs.pop('ref_mask', None)
if ref_mask is not None:
ref_mask = TF.resize(ref_mask, (h, w), antialias=True)
ref_mask = (ref_mask > 0.5).float()
ref_mask = rearrange(ref_mask, 'b c h w -> b (h w) c').contiguous()
x_in = x
x = self.norm(x)
if not self.use_linear:
x = self.proj_in(x)
x = rearrange(x, 'b c h w -> b (h w) c').contiguous()
if self.use_linear:
x = self.proj_in(x)
for i, block in enumerate(self.transformer_blocks):
if i > 0 and len(context) == 1:
i = 0 # use same context for each block
x = block(x, context=context[i], ref_mask=ref_mask, **kwargs)
if self.use_linear:
x = self.proj_out(x)
x = rearrange(x, 'b (h w) c -> b c h w', h=h, w=w).contiguous()
if not self.use_linear:
x = self.proj_out(x)
return x + x_in
@@ -1,17 +1,5 @@
# -*- coding: utf-8 -*-
# Copyright (c) Alibaba, Inc. and its affiliates.
import math
import torch
import torch.nn as nn
import torch.nn.functional
from einops import rearrange
from scepter.modules.model.backbone.video.init_helper import (
_init_transformer_weights, trunc_normal_)
from scepter.modules.model.registry import BACKBONES, BRICKS, STEMS
from scepter.modules.utils.config import dict_to_yaml
'''
The implementations of vivit as https://arxiv.org/abs/2103.15691.
The following setting alined the proposed model in the paper above.
@@ -39,6 +27,18 @@ TimesFormer:
complexity: (n_h * n_w) ** 2 + O(attn_temp)
'''
import math
import torch
import torch.nn as nn
import torch.nn.functional
from einops import rearrange
from scepter.modules.model.backbone.video.init_helper import (
_init_transformer_weights, trunc_normal_)
from scepter.modules.model.registry import BACKBONES, BRICKS, STEMS
from scepter.modules.utils.config import dict_to_yaml
@BACKBONES.register_class()
class VideoTransformer(nn.Module):
+4 -1
View File
@@ -1,7 +1,10 @@
# -*- coding: utf-8 -*-
# Copyright (c) Alibaba, Inc. and its affiliates.
from scepter.modules.model.embedder.embedder import (ConcatTimestepEmbedderND,
FrozenCLIPEmbedder,
FrozenOpenCLIPEmbedder,
FrozenOpenCLIPEmbedder2,
GeneralConditioner)
GeneralConditioner,
IPAdapterPlusEmbedder,
RefCrossEmbedder)
File diff suppressed because it is too large Load Diff
+137 -31
View File
@@ -22,9 +22,10 @@ from scepter.modules.utils.distribute import we
from scepter.modules.utils.file_system import FS
from .base_embedder import BaseEmbedder
from .resampler import Resampler
try:
from transformers import CLIPTextModel, CLIPTokenizer
from transformers import CLIPTextModel, CLIPTokenizer, CLIPVisionModelWithProjection
except Exception as e:
warnings.warn(
f'Import transformers error, please deal with this problem: {e}')
@@ -513,6 +514,93 @@ class ConcatTimestepEmbedderND(BaseEmbedder):
set_name=True)
@EMBEDDERS.register_class()
class IPAdapterPlusEmbedder(BaseEmbedder):
def __init__(self, cfg, logger=None):
super().__init__(cfg, logger=logger)
with FS.get_dir_to_local_dir(cfg.CLIP_DIR,
wait_finish=True) as local_path:
self.image_encoder = CLIPVisionModelWithProjection.from_pretrained(
local_path)
self.image_proj_model = Resampler(
dim=self.cfg.get('IN_DIM', 768),
depth=self.cfg.get('DEPTH', 4),
dim_head=64,
heads=self.cfg.get('HEADS', 12),
num_queries=self.cfg.get('NUM_TOKENS', 16),
embedding_dim=self.image_encoder.config.hidden_size,
output_dim=self.cfg.get('CROSSATTN_DIM', 768),
ff_mult=4,
)
with FS.get_from(cfg.PRETRAINED_MODEL, wait_finish=True) as local_path:
ckpt = torch.load(local_path, map_location='cpu')
self.image_proj_model.load_state_dict(ckpt['image_proj'],
strict=True)
self.patch_projector = nn.Linear(self.image_encoder.config.hidden_size,
self.cfg.get('CROSSATTN_DIM', 768))
def encode(self, ref_ip, ref_detail):
encoder_output = self.image_encoder(ref_ip, output_hidden_states=True)
image_prompt_embeds = self.image_proj_model(
encoder_output.hidden_states[-2])
encoder_output_2 = self.image_encoder(ref_detail,
output_hidden_states=True)
image_patch_embeds = self.patch_projector(
encoder_output_2.last_hidden_state)
out = {
'img_crossattn': image_prompt_embeds,
'ref_crossattn': image_patch_embeds,
}
return out
def forward(self, ref_ip, ref_detail):
return self.encode(ref_ip, ref_detail)
class RefCrossEmbedder(BaseEmbedder):
def __init__(self, cfg, logger=None):
super().__init__(cfg, logger=logger)
with FS.get_dir_to_local_dir(cfg.CLIP_DIR,
wait_finish=True) as local_path:
self.image_encoder = CLIPVisionModelWithProjection.from_pretrained(
local_path)
self.patch_projector = nn.Linear(self.image_encoder.config.hidden_size,
self.cfg.get('CROSSATTN_DIM', 768))
def encode(self, img):
encoder_output = self.image_encoder(img, output_hidden_states=True)
image_patch_embeds = self.patch_projector(
encoder_output.last_hidden_state)
out = {
'ref_crossattn': image_patch_embeds,
}
return out
def forward(self, img):
return self.encode(img)
@EMBEDDERS.register_class()
class TransparentEmbedder(BaseEmbedder):
def forward(self, *args):
out = dict()
for key, val in zip(self.input_keys, args):
out[key] = val
return out
@EMBEDDERS.register_class()
class NoiseConcatEmbedder(BaseEmbedder):
def forward(self, *args):
return {'concat': torch.cat(args, dim=1)}
@EMBEDDERS.register_class()
class GeneralConditioner(BaseEmbedder):
OUTPUT_DIM2KEYS = {2: 'y', 3: 'crossattn', 4: 'concat', 5: 'concat'}
@@ -598,42 +686,60 @@ class GeneralConditioner(BaseEmbedder):
with embedding_context():
if hasattr(embedder, 'input_key') and (embedder.input_key
is not None):
if embedder.input_key not in batch:
continue
if embedder.legacy_ucg_val is not None:
batch = self.possibly_get_ucg_val(embedder, batch)
emb_out = embedder(batch[embedder.input_key])
elif hasattr(embedder, 'input_keys'):
if any([k not in batch for k in embedder.input_keys]):
continue
emb_out = embedder(
*[batch[k] for k in embedder.input_keys])
assert isinstance(
emb_out, (torch.Tensor, list, tuple)
), f'encoder outputs must be tensors or a sequence, but got {type(emb_out)}'
if not isinstance(emb_out, (list, tuple)):
emb_out = [emb_out]
for emb in emb_out:
# print("emb.shape", emb.shape)
# print("emb.input_keys", embedder.input_keys)
out_key = self.OUTPUT_DIM2KEYS[emb.dim()]
if embedder.ucg_rate > 0.0 and embedder.legacy_ucg_val is None:
emb = (expand_dims_like(
torch.bernoulli(
(1.0 - embedder.ucg_rate) *
torch.ones(emb.shape[0], device=emb.device)),
emb,
) * emb)
if (hasattr(embedder, 'input_keys')):
if np.sum(
np.array([
key in force_zero_embeddings
for key in embedder.input_keys
])) > 0:
emb = torch.zeros_like(emb)
if out_key in output:
output[out_key] = torch.cat((output[out_key], emb),
self.KEY2CATDIM[out_key])
else:
output[out_key] = emb
# if "y" in output:
# print("out.shape", output["y"].shape)
if isinstance(emb_out, dict):
for key, val in emb_out.items():
if key in output:
# 重复出现的key必须在(y, crossattn, concat)中,否则raise error
assert key in self.KEY2CATDIM
output[key] = torch.cat([output[key], val],
dim=self.KEY2CATDIM[key])
else:
output[key] = val
else:
assert isinstance(
emb_out, (torch.Tensor, list, tuple)
), f'encoder outputs must be tensors or a sequence, but got {type(emb_out)}'
if not isinstance(emb_out, (list, tuple)):
emb_out = [emb_out]
for emb in emb_out:
# 根据emb的维度,判断该cond归属于 (y, concat, crossattn)中的哪一种
out_key = self.OUTPUT_DIM2KEYS[emb.dim()]
if embedder.ucg_rate > 0.0 and embedder.legacy_ucg_val is None:
emb = (expand_dims_like(
torch.bernoulli(
(1.0 - embedder.ucg_rate) *
torch.ones(emb.shape[0], device=emb.device)),
emb,
) * emb)
if (hasattr(embedder, 'input_keys')):
if np.sum(
np.array([
key in force_zero_embeddings
for key in embedder.input_keys
])) > 0:
emb = torch.zeros_like(emb)
if out_key in output:
output[out_key] = torch.cat((output[out_key], emb),
self.KEY2CATDIM[out_key])
else:
output[out_key] = emb
return output
def get_unconditional_conditioning(self,
+160
View File
@@ -0,0 +1,160 @@
# -*- coding: utf-8 -*-
# Copyright (c) Alibaba, Inc. and its affiliates.
import math
import torch
import torch.nn as nn
from einops import rearrange
from einops.layers.torch import Rearrange
# FFN
def FeedForward(dim, mult=4):
inner_dim = int(dim * mult)
return nn.Sequential(
nn.LayerNorm(dim),
nn.Linear(dim, inner_dim, bias=False),
nn.GELU(),
nn.Linear(inner_dim, dim, bias=False),
)
def reshape_tensor(x, heads):
bs, length, width = x.shape
# (bs, length, width) --> (bs, length, n_heads, dim_per_head)
x = x.view(bs, length, heads, -1)
# (bs, length, n_heads, dim_per_head) --> (bs, n_heads, length, dim_per_head)
x = x.transpose(1, 2)
# (bs, n_heads, length, dim_per_head) --> (bs*n_heads, length, dim_per_head)
x = x.reshape(bs, heads, length, -1)
return x
class PerceiverAttention(nn.Module):
def __init__(self, *, dim, dim_head=64, heads=8):
super().__init__()
self.scale = dim_head**-0.5
self.dim_head = dim_head
self.heads = heads
inner_dim = dim_head * heads
self.norm1 = nn.LayerNorm(dim)
self.norm2 = nn.LayerNorm(dim)
self.to_q = nn.Linear(dim, inner_dim, bias=False)
self.to_kv = nn.Linear(dim, inner_dim * 2, bias=False)
self.to_out = nn.Linear(inner_dim, dim, bias=False)
def forward(self, x, latents):
"""
Args:
x (torch.Tensor): image features
shape (b, n1, D)
latent (torch.Tensor): latent features
shape (b, n2, D)
"""
x = self.norm1(x)
latents = self.norm2(latents)
b, l, _ = latents.shape
q = self.to_q(latents)
kv_input = torch.cat((x, latents), dim=-2)
k, v = self.to_kv(kv_input).chunk(2, dim=-1)
q = reshape_tensor(q, self.heads)
k = reshape_tensor(k, self.heads)
v = reshape_tensor(v, self.heads)
# attention
scale = 1 / math.sqrt(math.sqrt(self.dim_head))
weight = (q * scale) @ (k * scale).transpose(
-2, -1) # More stable with f16 than dividing afterwards
weight = torch.softmax(weight.float(), dim=-1).type(weight.dtype)
out = weight @ v
out = out.permute(0, 2, 1, 3).reshape(b, l, -1)
return self.to_out(out)
class Resampler(nn.Module):
def __init__(
self,
dim=1024,
depth=8,
dim_head=64,
heads=16,
num_queries=8,
embedding_dim=768,
output_dim=1024,
ff_mult=4,
max_seq_len: int = 257, # CLIP tokens + CLS token
apply_pos_emb: bool = False,
num_latents_mean_pooled:
int = 0, # number of latents derived from mean pooled representation of the sequence
):
super().__init__()
self.pos_emb = nn.Embedding(max_seq_len,
embedding_dim) if apply_pos_emb else None
self.latents = nn.Parameter(
torch.randn(1, num_queries, dim) / dim**0.5)
self.proj_in = nn.Linear(embedding_dim, dim)
self.proj_out = nn.Linear(dim, output_dim)
self.norm_out = nn.LayerNorm(output_dim)
self.to_latents_from_mean_pooled_seq = (nn.Sequential(
nn.LayerNorm(dim),
nn.Linear(dim, dim * num_latents_mean_pooled),
Rearrange('b (n d) -> b n d', n=num_latents_mean_pooled),
) if num_latents_mean_pooled > 0 else None)
self.layers = nn.ModuleList([])
for _ in range(depth):
self.layers.append(
nn.ModuleList([
PerceiverAttention(dim=dim, dim_head=dim_head,
heads=heads),
FeedForward(dim=dim, mult=ff_mult),
]))
def forward(self, x):
if self.pos_emb is not None:
n, device = x.shape[1], x.device
pos_emb = self.pos_emb(torch.arange(n, device=device))
x = x + pos_emb
latents = self.latents.repeat(x.size(0), 1, 1)
x = self.proj_in(x)
if self.to_latents_from_mean_pooled_seq:
meanpooled_seq = masked_mean(x,
dim=1,
mask=torch.ones(x.shape[:2],
device=x.device,
dtype=torch.bool))
meanpooled_latents = self.to_latents_from_mean_pooled_seq(
meanpooled_seq)
latents = torch.cat((meanpooled_latents, latents), dim=-2)
for attn, ff in self.layers:
latents = attn(x, latents) + latents
latents = ff(latents) + latents
latents = self.proj_out(latents)
return self.norm_out(latents)
def masked_mean(t, *, dim, mask=None):
if mask is None:
return t.mean(dim=dim)
denom = mask.sum(dim=dim, keepdim=True)
mask = rearrange(mask, 'b n -> b n 1')
masked_t = t.masked_fill(~mask, 0.0)
return masked_t.sum(dim=dim) / denom.clamp(min=1e-5)
+6 -3
View File
@@ -1,5 +1,8 @@
# -*- coding: utf-8 -*-
# Copyright (c) Alibaba, Inc. and its affiliates.
from scepter.modules.model.head.classifier_head import (
ClassifierHead, CosineLinearHead, TransformerHead, TransformerHeadx2,
VideoClassifierHead, VideoClassifierHeadx2)
from scepter.modules.model.head.classifier_head import (ClassifierHead,
CosineLinearHead,
TransformerHead,
TransformerHeadx2,
VideoClassifierHead,
VideoClassifierHeadx2)
+2 -3
View File
@@ -1,5 +1,4 @@
# -*- coding: utf-8 -*-
# Copyright (c) Alibaba, Inc. and its affiliates.
from scepter.modules.model.metric.classification import (AccuracyMetric,
EnsembleAccuracyMetric
)
from scepter.modules.model.metric.classification import (
AccuracyMetric, EnsembleAccuracyMetric)
@@ -13,8 +13,7 @@ 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)
sample_euler, sample_euler_ancestral, sample_heun)
__all__ = ['GaussianDiffusion']
@@ -27,6 +26,148 @@ def _i(tensor, t, x):
return tensor[t.to(tensor.device)].view(shape).to(x.device)
def _unpack_2d_ks(kernel_size):
if isinstance(kernel_size, int):
ky = kx = kernel_size
else:
assert len(
kernel_size) == 2, '2D Kernel size should have a length of 2.'
ky, kx = kernel_size
ky = int(ky)
kx = int(kx)
return ky, kx
def _compute_zero_padding(kernel_size):
ky, kx = _unpack_2d_ks(kernel_size)
return (ky - 1) // 2, (kx - 1) // 2
def _bilateral_blur(
input,
guidance,
kernel_size,
sigma_color,
sigma_space,
border_type='reflect',
color_distance_type='l1',
):
if isinstance(sigma_color, torch.Tensor):
sigma_color = sigma_color.to(device=input.device,
dtype=input.dtype).view(-1, 1, 1, 1, 1)
ky, kx = _unpack_2d_ks(kernel_size)
pad_y, pad_x = _compute_zero_padding(kernel_size)
padded_input = torch.nn.functional.pad(input, (pad_x, pad_x, pad_y, pad_y),
mode=border_type)
unfolded_input = padded_input.unfold(2, ky, 1).unfold(3, kx, 1).flatten(
-2) # (B, C, H, W, Ky x Kx)
if guidance is None:
guidance = input
unfolded_guidance = unfolded_input
else:
padded_guidance = torch.nn.functional.pad(guidance,
(pad_x, pad_x, pad_y, pad_y),
mode=border_type)
unfolded_guidance = padded_guidance.unfold(2, ky, 1).unfold(
3, kx, 1).flatten(-2) # (B, C, H, W, Ky x Kx)
diff = unfolded_guidance - guidance.unsqueeze(-1)
if color_distance_type == 'l1':
color_distance_sq = diff.abs().sum(1, keepdim=True).square()
elif color_distance_type == 'l2':
color_distance_sq = diff.square().sum(1, keepdim=True)
else:
raise ValueError('color_distance_type only acceps l1 or l2')
color_kernel = (-0.5 / sigma_color**2 *
color_distance_sq).exp() # (B, 1, H, W, Ky x Kx)
space_kernel = get_gaussian_kernel2d(kernel_size,
sigma_space,
device=input.device,
dtype=input.dtype)
space_kernel = space_kernel.view(-1, 1, 1, 1, kx * ky)
kernel = space_kernel * color_kernel
out = (unfolded_input * kernel).sum(-1) / kernel.sum(-1)
return out
def get_gaussian_kernel1d(
kernel_size,
sigma,
force_even,
*,
device=None,
dtype=None,
):
return gaussian(kernel_size, sigma, device=device, dtype=dtype)
def gaussian(window_size, sigma, *, device=None, dtype=None):
batch_size = sigma.shape[0]
x = (torch.arange(window_size, device=sigma.device, dtype=sigma.dtype) -
window_size // 2).expand(batch_size, -1)
if window_size % 2 == 0:
x = x + 0.5
gauss = torch.exp(-x.pow(2.0) / (2 * sigma.pow(2.0)))
return gauss / gauss.sum(-1, keepdim=True)
def get_gaussian_kernel2d(
kernel_size,
sigma,
force_even=False,
*,
device=None,
dtype=None,
):
sigma = torch.Tensor([[sigma, sigma]]).to(device=device, dtype=dtype)
ksize_y, ksize_x = _unpack_2d_ks(kernel_size)
sigma_y, sigma_x = sigma[:, 0, None], sigma[:, 1, None]
kernel_y = get_gaussian_kernel1d(ksize_y,
sigma_y,
force_even,
device=device,
dtype=dtype)[..., None]
kernel_x = get_gaussian_kernel1d(ksize_x,
sigma_x,
force_even,
device=device,
dtype=dtype)[..., None]
return kernel_y * kernel_x.view(-1, 1, ksize_x)
def adaptive_anisotropic_filter(x, g=None):
if g is None:
g = x
s, m = torch.std_mean(g, dim=(1, 2, 3), keepdim=True)
s = s + 1e-5
guidance = (g - m) / s
y = _bilateral_blur(x,
guidance,
kernel_size=(13, 13),
sigma_color=3.0,
sigma_space=3.0,
border_type='reflect',
color_distance_type='l1')
return y
class GaussianDiffusion(object):
def __init__(self, sigmas, prediction_type='eps'):
assert prediction_type in {'x0', 'eps', 'v'}
@@ -53,6 +194,7 @@ class GaussianDiffusion(object):
guide_scale=None,
guide_rescale=None,
clamp=None,
sharpness=0.0,
percentile=None,
cat_uc=False,
**kwargs):
@@ -79,54 +221,99 @@ class GaussianDiffusion(object):
# prediction
if guide_scale is None:
assert isinstance(model_kwargs, dict)
out = model(xt, t=t, **model_kwargs, **kwargs)
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 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
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)
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:
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 cat_uc:
# 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
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,
**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)
# todo sharpness
# sharpness sampling
if sharpness is not None and sharpness > 0:
positive_x0 = alphas * xt - sigmas * y_out
negative_x0 = alphas * xt - sigmas * u_out
positive_eps = xt - positive_x0
negative_eps = xt - negative_x0
global_diffusion_progress = (
1 - t / 999.0).detach().cpu().numpy().tolist()[0]
alpha = 0.001 * sharpness * global_diffusion_progress
positive_eps_degraded = adaptive_anisotropic_filter(
x=positive_eps, g=positive_x0)
positive_eps_degraded_weighted = positive_eps_degraded * alpha + positive_eps * (
1.0 - alpha)
final_eps = negative_eps + guide_scale * (
positive_eps_degraded_weighted - negative_eps)
final_x0 = xt - final_eps
out = (alphas * xt - final_x0) / sigmas
else:
out = u_out + guide_scale * (y_out - u_out)
elif isinstance(guide_scale, dict):
assert len(model_kwargs) == 3
y_out = model(xt, t=t, **model_kwargs[0], **kwargs)
m_out = model(xt, t=t, **model_kwargs[1], **kwargs)
u_out = model(xt, t=t, **model_kwargs[2], **kwargs)
out = u_out + guide_scale['image'] * (
m_out - u_out) + guide_scale['text'] * (y_out - m_out)
elif isinstance(guide_scale, list):
assert len(guide_scale) == len(model_kwargs) - 1
y_out = model(xt, t=t, **model_kwargs[0], **kwargs)
outs = [y_out]
for i in range(1, len(model_kwargs)):
outs.append(model(xt, t=t, **model_kwargs[i], **kwargs))
out = outs[-1]
for i in range(len(guide_scale)):
out += guide_scale[i] * (outs[-i - 2] - outs[-i - 1])
# rescale the output according to arXiv:2305.08891
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
# compute x0
if self.prediction_type == 'x0':
x0 = out
@@ -197,6 +384,7 @@ class GaussianDiffusion(object):
guide_scale=None,
guide_rescale=None,
clamp=None,
sharpness=0.0,
percentile=None,
solver='euler_a',
steps=20,
@@ -209,12 +397,16 @@ class GaussianDiffusion(object):
seed=-1,
intermediate_callback=None,
cat_uc=False,
add_noise=False,
free_steps=None,
step_offset=None,
**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 discretization in (None, 'leading', 'linspace', 'trailing',
'free')
assert discard_penultimate_step in (None, True, False)
assert return_intermediate in (None, 'x0', 'xt')
@@ -255,17 +447,50 @@ class GaussianDiffusion(object):
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,
**kwargs)[-2]
if isinstance(
model_kwargs[0]['cond'], dict) and \
'tar_x0' in model_kwargs[0]['cond'] and \
'tar_mask_latent' in model_kwargs[0]['cond']:
tar_x0 = model_kwargs[0]['cond']['tar_x0']
tar_mask = model_kwargs[0]['cond']['tar_mask_latent']
tar_xt = self.diffuse(x0=tar_x0, t=t)
xt = tar_xt * (1.0 - tar_mask) + xt * tar_mask
if isinstance(model_kwargs[0]['cond'],
dict) and 'ref_x0' in model_kwargs[0]['cond']:
model_kwargs[0]['cond']['ref_xt'] = self.diffuse(
x0=model_kwargs[0]['cond']['ref_x0'], t=t)
model_kwargs[1]['cond']['ref_xt'] = self.diffuse(
x0=model_kwargs[1]['cond']['ref_x0'], t=t)
if solver in ('onestep', 'multistep', 'multistep2', 'multistep3'):
x0 = self.denoise(xt,
t,
None,
model,
model_kwargs,
guide_scale,
guide_rescale,
clamp,
sharpness,
percentile,
cat_uc=cat_uc,
**kwargs)[-3]
else:
x0 = self.denoise(xt,
t,
None,
model,
model_kwargs,
guide_scale,
guide_rescale,
clamp,
sharpness,
percentile,
cat_uc=cat_uc,
**kwargs)[-2]
# collect intermediate outputs
if return_intermediate == 'xt':
@@ -291,10 +516,14 @@ class GaussianDiffusion(object):
elif discretization == 'trailing':
steps = torch.arange(t_max, t_min - 1,
-((t_max - t_min + 1) / steps))
elif discretization == 'free':
steps = torch.tensor(free_steps)
else:
raise NotImplementedError(
f'{discretization} discretization not implemented')
steps = steps.clamp_(t_min, t_max)
elif isinstance(steps, list):
steps = torch.tensor(steps)
steps = torch.as_tensor(steps,
dtype=torch.float32,
device=noise.device)
@@ -335,6 +564,23 @@ class GaussianDiffusion(object):
if discard_penultimate_step:
sigmas = torch.cat([sigmas[:-2], sigmas[-1:]])
kwargs['seed'] = seed
# add noise to x0
if add_noise:
if 'dm_steps' in kwargs:
if step_offset:
add_noise_step = -kwargs['dm_steps'] + step_offset
if add_noise_step < 0:
noise = self.diffuse(
noise,
torch.full((noise.shape[0], 1),
steps[add_noise_step],
dtype=torch.int))
else:
noise = self.diffuse(
noise,
torch.full((noise.shape[0], 1),
steps[-kwargs['dm_steps'] - 1],
dtype=torch.int))
# sampling
x0 = solver_fn(noise,
model_fn,
@@ -373,168 +619,6 @@ class GaussianDiffusion(object):
| 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,
**kwargs)[-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
+148
View File
@@ -0,0 +1,148 @@
# -*- coding: utf-8 -*-
# Copyright (c) Alibaba, Inc. and its affiliates.
import cv2
import numpy as np
import torchvision.transforms.functional as TF
def get_bbox_from_mask(mask):
h, w = mask.shape[0], mask.shape[1]
if mask.sum() < 10:
return 0, h, 0, w
rows = np.any(mask, axis=1)
cols = np.any(mask, axis=0)
y1, y2 = np.where(rows)[0][[0, -1]]
x1, x2 = np.where(cols)[0][[0, -1]]
return (y1, y2, x1, x2)
def pad_to_square(image, pad_value=255, random=False):
H, W = image.shape[0], image.shape[1]
if H == W:
return image, 0, 0
padd = abs(H - W)
if random:
padd_1 = int(np.random.randint(0, padd))
else:
padd_1 = int(padd / 2)
padd_2 = padd - padd_1
if H > W:
pad_param = ((0, 0), (padd_1, padd_2), (0, 0))
else:
pad_param = ((padd_1, padd_2), (0, 0), (0, 0))
# print(pad_param, pad_value)
image = np.pad(image, pad_param, 'constant', constant_values=pad_value)
return image, padd_1, padd_2
def box_in_box(small_box, big_box):
y1, y2, x1, x2 = small_box
y1_b, _, x1_b, _ = big_box
y1, y2, x1, x2 = y1 - y1_b, y2 - y1_b, x1 - x1_b, x2 - x1_b
return (y1, y2, x1, x2)
def box2squre(image, box):
H, W = image.shape[0], image.shape[1]
y1, y2, x1, x2 = box
cx = (x1 + x2) // 2
cy = (y1 + y2) // 2
h, w = y2 - y1, x2 - x1
if h >= w:
x1 = cx - h // 2
x2 = x1 + h
else:
y1 = cy - w // 2
y2 = y1 + w
x1 = max(0, x1)
x2 = min(W, x2)
y1 = max(0, y1)
y2 = min(H, y2)
return (y1, y2, x1, x2)
def expand_bbox(mask,
yyxx,
ratio=1.0,
min_crop=0,
expand_type='center',
to_square=False):
y1, y2, x1, x2 = yyxx
h = y2 - y1 + 1
w = x2 - x1 + 1
H, W = mask.shape[0], mask.shape[1]
xc, yc = 0.5 * (x1 + x2), 0.5 * (y1 + y2)
def expand(k):
if isinstance(ratio, tuple) or isinstance(ratio, list):
r = np.random.uniform(*ratio)
k = k * r
else:
k = ratio * k
return k
new_h = expand(h)
new_w = expand(w)
new_h = max(new_h, min_crop)
new_w = max(new_w, min_crop)
if to_square:
if new_w / new_h < 0.334:
new_w = new_w + 1.0 / 3.0 * new_h
elif new_h / new_w < 0.334:
new_h = new_h + 1.0 / 3.0 * new_w
if expand_type == 'center':
x1 = max(0, int(xc - new_w * 0.5))
x2 = min(W, int(xc + new_w * 0.5))
y1 = max(0, int(yc - new_h * 0.5))
y2 = min(H, int(yc + new_h * 0.5))
else:
x1 = max(0, min(x1,
int(x2 - new_w * np.random.uniform(w / new_w, 1.0))))
x2 = min(W, max(x2, x1 + new_w))
y1 = max(0, min(y1,
int(y2 - new_h * np.random.uniform(h / new_h, 1.0))))
y2 = min(H, max(y2, y1 + new_h))
return (int(y1), int(y2), int(x1), int(x2))
def crop_back(pred, tar_image, extra_sizes, tar_box_yyxx_crop):
H1, W1, H2, W2, pad1, pad2 = extra_sizes
y1, y2, x1, x2 = tar_box_yyxx_crop
pred = TF.resize(pred, (H2, W2), antialias=True)
# if W1 == H1:
# tar_image[:, y1:y2, x1:x2] = pred
# return tar_image
# if W1 < W2:
# pad1 = int((W2 - W1) / 2)
# pad2 = W2 - W1 - pad1
# pred = pred[:, :,pad1:-pad2]
# else:
# pad1 = int((H2 - H1) / 2)
# pad2 = H2 - H1 - pad1
# pred = pred[:, pad1:-pad2, :]
if W1 < W2:
# pad width
assert H1 == H2 and (pad1 + W1) == (W2 - pad2)
pred = pred[:, :, pad1 + 2:(W2 - pad2 - 2)]
tar_image[:, y1:y2, x1 + 2:x2 - 2] = pred
elif H1 < H2:
# pad height
assert W1 == W2 and (pad1 + H1) == (H2 - pad2)
pred = pred[:, pad1 + 2:(H2 - pad2 - 2), :]
tar_image[:, y1 + 2:y2 - 2, x1:x2] = pred
else:
tar_image[:, y1:y2, x1:x2] = pred
return tar_image
def save_image(image, save_path):
image = cv2.cvtColor(image, cv2.COLOR_RGB2BGR)
cv2.imwrite(save_path, image)