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