initial commit
@@ -0,0 +1,146 @@
|
|||||||
|
.idea/
|
||||||
|
|
||||||
|
training/
|
||||||
|
lightning_logs/
|
||||||
|
image_log/
|
||||||
|
|
||||||
|
*.pth
|
||||||
|
*.pt
|
||||||
|
*.ckpt
|
||||||
|
*.safetensors
|
||||||
|
|
||||||
|
# Byte-compiled / optimized / DLL files
|
||||||
|
__pycache__/
|
||||||
|
*.py[cod]
|
||||||
|
*$py.class
|
||||||
|
|
||||||
|
# C extensions
|
||||||
|
*.so
|
||||||
|
|
||||||
|
# Distribution / packaging
|
||||||
|
.Python
|
||||||
|
build/
|
||||||
|
develop-eggs/
|
||||||
|
dist/
|
||||||
|
downloads/
|
||||||
|
eggs/
|
||||||
|
.eggs/
|
||||||
|
lib/
|
||||||
|
lib64/
|
||||||
|
parts/
|
||||||
|
sdist/
|
||||||
|
var/
|
||||||
|
wheels/
|
||||||
|
pip-wheel-metadata/
|
||||||
|
share/python-wheels/
|
||||||
|
*.egg-info/
|
||||||
|
.installed.cfg
|
||||||
|
*.egg
|
||||||
|
MANIFEST
|
||||||
|
|
||||||
|
# PyInstaller
|
||||||
|
# Usually these files are written by a python script from a template
|
||||||
|
# before PyInstaller builds the exe, so as to inject date/other infos into it.
|
||||||
|
*.manifest
|
||||||
|
*.spec
|
||||||
|
|
||||||
|
# Installer logs
|
||||||
|
pip-log.txt
|
||||||
|
pip-delete-this-directory.txt
|
||||||
|
|
||||||
|
# Unit test / coverage reports
|
||||||
|
htmlcov/
|
||||||
|
.tox/
|
||||||
|
.nox/
|
||||||
|
.coverage
|
||||||
|
.coverage.*
|
||||||
|
.cache
|
||||||
|
nosetests.xml
|
||||||
|
coverage.xml
|
||||||
|
*.cover
|
||||||
|
*.py,cover
|
||||||
|
.hypothesis/
|
||||||
|
.pytest_cache/
|
||||||
|
|
||||||
|
# Translations
|
||||||
|
*.mo
|
||||||
|
*.pot
|
||||||
|
|
||||||
|
# Django stuff:
|
||||||
|
*.log
|
||||||
|
local_settings.py
|
||||||
|
db.sqlite3
|
||||||
|
db.sqlite3-journal
|
||||||
|
|
||||||
|
# Flask stuff:
|
||||||
|
instance/
|
||||||
|
.webassets-cache
|
||||||
|
|
||||||
|
# Scrapy stuff:
|
||||||
|
.scrapy
|
||||||
|
|
||||||
|
# Sphinx documentation
|
||||||
|
docs/_build/
|
||||||
|
|
||||||
|
# PyBuilder
|
||||||
|
target/
|
||||||
|
|
||||||
|
# Jupyter Notebook
|
||||||
|
.ipynb_checkpoints
|
||||||
|
|
||||||
|
# IPython
|
||||||
|
profile_default/
|
||||||
|
ipython_config.py
|
||||||
|
|
||||||
|
# pyenv
|
||||||
|
.python-version
|
||||||
|
|
||||||
|
# pipenv
|
||||||
|
# According to pypa/pipenv#598, it is recommended to include Pipfile.lock in version control.
|
||||||
|
# However, in case of collaboration, if having platform-specific dependencies or dependencies
|
||||||
|
# having no cross-platform support, pipenv may install dependencies that don't work, or not
|
||||||
|
# install all needed dependencies.
|
||||||
|
#Pipfile.lock
|
||||||
|
|
||||||
|
# PEP 582; used by e.g. github.com/David-OConnor/pyflow
|
||||||
|
__pypackages__/
|
||||||
|
|
||||||
|
# Celery stuff
|
||||||
|
celerybeat-schedule
|
||||||
|
celerybeat.pid
|
||||||
|
|
||||||
|
# SageMath parsed files
|
||||||
|
*.sage.py
|
||||||
|
|
||||||
|
# Environments
|
||||||
|
.env
|
||||||
|
.venv
|
||||||
|
env/
|
||||||
|
venv/
|
||||||
|
ENV/
|
||||||
|
env.bak/
|
||||||
|
venv.bak/
|
||||||
|
|
||||||
|
# Spyder project settings
|
||||||
|
.spyderproject
|
||||||
|
.spyproject
|
||||||
|
|
||||||
|
# Rope project settings
|
||||||
|
.ropeproject
|
||||||
|
|
||||||
|
# mkdocs documentation
|
||||||
|
/site
|
||||||
|
|
||||||
|
# mypy
|
||||||
|
.mypy_cache/
|
||||||
|
.dmypy.json
|
||||||
|
dmypy.json
|
||||||
|
|
||||||
|
# Pyre type checker
|
||||||
|
.pyre/
|
||||||
|
|
||||||
|
|
||||||
|
*.safetensors
|
||||||
|
*.ckpt
|
||||||
|
|
||||||
|
checkpoints
|
||||||
@@ -0,0 +1,7 @@
|
|||||||
|
# ComfyUI wrapper nodes for LVCD:
|
||||||
|
|
||||||
|
|
||||||
|
|
||||||
|
Original repo:
|
||||||
|
|
||||||
|
https://github.com/luckyhzt/LVCD
|
||||||
@@ -0,0 +1,3 @@
|
|||||||
|
from .nodes import NODE_CLASS_MAPPINGS, NODE_DISPLAY_NAME_MAPPINGS
|
||||||
|
|
||||||
|
__all__ = ["NODE_CLASS_MAPPINGS", "NODE_DISPLAY_NAME_MAPPINGS"]
|
||||||
@@ -0,0 +1,250 @@
|
|||||||
|
model:
|
||||||
|
base_learning_rate: 5.0e-5
|
||||||
|
target: .models.csvd.VideoDiffusionEngine
|
||||||
|
params:
|
||||||
|
scale_factor: 0.18215
|
||||||
|
disable_first_stage_autocast: True
|
||||||
|
ckpt_path: checkpoints/svd.safetensors
|
||||||
|
control_model_path: Null
|
||||||
|
init_from_unet: True
|
||||||
|
sd_locked: False
|
||||||
|
drop_first_stage_model: True
|
||||||
|
|
||||||
|
denoiser_config:
|
||||||
|
target: .sgm.modules.diffusionmodules.denoiser.Denoiser
|
||||||
|
params:
|
||||||
|
scaling_config:
|
||||||
|
target: .sgm.modules.diffusionmodules.denoiser_scaling.VScalingWithEDMcNoise
|
||||||
|
|
||||||
|
network_config:
|
||||||
|
target: .models.csvd.ControlledVideoUNet
|
||||||
|
params:
|
||||||
|
adm_in_channels: 768
|
||||||
|
num_classes: sequential
|
||||||
|
use_checkpoint: True
|
||||||
|
in_channels: 8
|
||||||
|
out_channels: 4
|
||||||
|
model_channels: 320
|
||||||
|
attention_resolutions: [4, 2, 1]
|
||||||
|
num_res_blocks: 2
|
||||||
|
channel_mult: [1, 2, 4, 4]
|
||||||
|
num_head_channels: 64
|
||||||
|
use_linear_in_transformer: True
|
||||||
|
transformer_depth: 1
|
||||||
|
context_dim: 1024
|
||||||
|
spatial_transformer_attn_type: softmax-xformers
|
||||||
|
extra_ff_mix_layer: True
|
||||||
|
use_spatial_context: True
|
||||||
|
merge_strategy: learned_with_images
|
||||||
|
video_kernel_size: [3, 1, 1]
|
||||||
|
temporal_attn_type: .models.layers.TemporalAttention_Masked
|
||||||
|
spatial_self_attn_type: .models.layers.ReferenceAttention
|
||||||
|
conv3d_type: .models.layers.Conv3d_Masked
|
||||||
|
trainable_layers: ['TemporalAttention_Masked', 'ReferenceAttention']
|
||||||
|
|
||||||
|
controlnet_config:
|
||||||
|
target: .models.csvd.ControlNet
|
||||||
|
params:
|
||||||
|
adm_in_channels: 768
|
||||||
|
num_classes: sequential
|
||||||
|
use_checkpoint: True
|
||||||
|
in_channels: 8
|
||||||
|
model_channels: 320
|
||||||
|
hint_channels: 3
|
||||||
|
attention_resolutions: [4, 2, 1]
|
||||||
|
num_res_blocks: 2
|
||||||
|
channel_mult: [1, 2, 4, 4]
|
||||||
|
num_head_channels: 64
|
||||||
|
use_linear_in_transformer: True
|
||||||
|
transformer_depth: 1
|
||||||
|
context_dim: 1024
|
||||||
|
spatial_transformer_attn_type: softmax-xformers
|
||||||
|
extra_ff_mix_layer: True
|
||||||
|
use_spatial_context: True
|
||||||
|
merge_strategy: learned_with_images
|
||||||
|
video_kernel_size: [3, 1, 1]
|
||||||
|
temporal_attn_type: .models.layers.TemporalAttention_Masked
|
||||||
|
spatial_self_attn_type: .models.layers.ReferenceAttention
|
||||||
|
conv3d_type: .models.layers.Conv3d_Masked
|
||||||
|
|
||||||
|
conditioner_config:
|
||||||
|
target: .sgm.modules.GeneralConditioner
|
||||||
|
params:
|
||||||
|
emb_models:
|
||||||
|
- is_trainable: False
|
||||||
|
input_key: cond_frames_without_noise
|
||||||
|
target: .sgm.modules.encoders.modules.FrozenOpenCLIPImagePredictionEmbedder
|
||||||
|
params:
|
||||||
|
n_cond_frames: 1
|
||||||
|
n_copies: 1
|
||||||
|
open_clip_embedding_config:
|
||||||
|
target: .sgm.modules.encoders.modules.FrozenOpenCLIPImageEmbedder
|
||||||
|
params:
|
||||||
|
freeze: True
|
||||||
|
init_device : cuda:0
|
||||||
|
|
||||||
|
- input_key: fps_id
|
||||||
|
is_trainable: False
|
||||||
|
target: .sgm.modules.encoders.modules.ConcatTimestepEmbedderND
|
||||||
|
params:
|
||||||
|
outdim: 256
|
||||||
|
|
||||||
|
- input_key: motion_bucket_id
|
||||||
|
is_trainable: False
|
||||||
|
target: .sgm.modules.encoders.modules.ConcatTimestepEmbedderND
|
||||||
|
params:
|
||||||
|
outdim: 256
|
||||||
|
|
||||||
|
- input_key: cond_frames
|
||||||
|
is_trainable: False
|
||||||
|
target: .sgm.modules.encoders.modules.VideoPredictionEmbedderWithEncoder
|
||||||
|
params:
|
||||||
|
disable_encoder_autocast: True
|
||||||
|
n_cond_frames: 1
|
||||||
|
n_copies: 1
|
||||||
|
is_ae: True
|
||||||
|
encoder_config:
|
||||||
|
target: .sgm.models.autoencoder.AutoencoderKLModeOnly
|
||||||
|
params:
|
||||||
|
embed_dim: 4
|
||||||
|
monitor: val/rec_loss
|
||||||
|
ddconfig:
|
||||||
|
attn_type: vanilla-xformers
|
||||||
|
double_z: True
|
||||||
|
z_channels: 4
|
||||||
|
resolution: 256
|
||||||
|
in_channels: 3
|
||||||
|
out_ch: 3
|
||||||
|
ch: 128
|
||||||
|
ch_mult: [1, 2, 4, 4]
|
||||||
|
num_res_blocks: 2
|
||||||
|
attn_resolutions: []
|
||||||
|
dropout: 0.0
|
||||||
|
lossconfig:
|
||||||
|
target: torch.nn.Identity
|
||||||
|
|
||||||
|
- input_key: cond_aug
|
||||||
|
is_trainable: False
|
||||||
|
target: .sgm.modules.encoders.modules.ConcatTimestepEmbedderND
|
||||||
|
params:
|
||||||
|
outdim: 256
|
||||||
|
|
||||||
|
first_stage_config:
|
||||||
|
target: .sgm.models.autoencoder.AutoencodingEngine
|
||||||
|
params:
|
||||||
|
loss_config:
|
||||||
|
target: torch.nn.Identity
|
||||||
|
regularizer_config:
|
||||||
|
target: .sgm.modules.autoencoding.regularizers.DiagonalGaussianRegularizer
|
||||||
|
encoder_config:
|
||||||
|
target: .sgm.modules.diffusionmodules.model.Encoder
|
||||||
|
params:
|
||||||
|
attn_type: vanilla
|
||||||
|
double_z: True
|
||||||
|
z_channels: 4
|
||||||
|
resolution: 256
|
||||||
|
in_channels: 3
|
||||||
|
out_ch: 3
|
||||||
|
ch: 128
|
||||||
|
ch_mult: [1, 2, 4, 4]
|
||||||
|
num_res_blocks: 2
|
||||||
|
attn_resolutions: []
|
||||||
|
dropout: 0.0
|
||||||
|
decoder_config:
|
||||||
|
target: .sgm.modules.autoencoding.temporal_ae.VideoDecoder
|
||||||
|
params:
|
||||||
|
attn_type: vanilla
|
||||||
|
double_z: True
|
||||||
|
z_channels: 4
|
||||||
|
resolution: 256
|
||||||
|
in_channels: 3
|
||||||
|
out_ch: 3
|
||||||
|
ch: 128
|
||||||
|
ch_mult: [1, 2, 4, 4]
|
||||||
|
num_res_blocks: 2
|
||||||
|
attn_resolutions: []
|
||||||
|
dropout: 0.0
|
||||||
|
video_kernel_size: [3, 1, 1]
|
||||||
|
|
||||||
|
sampler_config:
|
||||||
|
target: .sgm.modules.diffusionmodules.sampling.EulerEDMSampler
|
||||||
|
params:
|
||||||
|
num_steps: 25
|
||||||
|
|
||||||
|
discretization_config:
|
||||||
|
target: .sgm.modules.diffusionmodules.discretizer.EDMDiscretization
|
||||||
|
params:
|
||||||
|
sigma_max: 700.0
|
||||||
|
|
||||||
|
guider_config:
|
||||||
|
target: .sgm.modules.diffusionmodules.guiders.LinearPredictionGuider
|
||||||
|
params:
|
||||||
|
num_frames: 14
|
||||||
|
max_scale: 2.5
|
||||||
|
min_scale: 1.0
|
||||||
|
additional_cond_keys: ['control_hint']
|
||||||
|
|
||||||
|
loss_fn_config:
|
||||||
|
target: .sgm.modules.diffusionmodules.loss.StandardDiffusionLoss
|
||||||
|
params:
|
||||||
|
batch2model_keys: ['num_video_frames', 'image_only_indicator']
|
||||||
|
additional_cond_keys: ['control_hint', 'crossattn_scale', 'concat_scale']
|
||||||
|
|
||||||
|
loss_weighting_config:
|
||||||
|
target: .sgm.modules.diffusionmodules.loss_weighting.EDMWeighting
|
||||||
|
params:
|
||||||
|
sigma_data: 1.0
|
||||||
|
|
||||||
|
sigma_sampler_config:
|
||||||
|
target: .sgm.modules.diffusionmodules.sigma_sampling.EDMSampling
|
||||||
|
params:
|
||||||
|
p_mean: 1.0
|
||||||
|
p_std: 1.6
|
||||||
|
|
||||||
|
|
||||||
|
lightning:
|
||||||
|
modelcheckpoint:
|
||||||
|
params:
|
||||||
|
every_n_train_steps: 1500
|
||||||
|
save_last: False
|
||||||
|
save_top_k: -1
|
||||||
|
filename: '{epoch:04d}-{global_step:06.0f}'
|
||||||
|
|
||||||
|
strategy:
|
||||||
|
params:
|
||||||
|
process_group_backend: gloo
|
||||||
|
|
||||||
|
trainer:
|
||||||
|
devices: 4,5,6,7,
|
||||||
|
benchmark: True
|
||||||
|
num_sanity_val_steps: 0
|
||||||
|
accumulate_grad_batches: 4
|
||||||
|
max_epochs: 100
|
||||||
|
precision: 16-mixed
|
||||||
|
|
||||||
|
|
||||||
|
data:
|
||||||
|
target: .sgm.data.my_dataset.DataModuleFromConfig
|
||||||
|
params:
|
||||||
|
batch_size: 2
|
||||||
|
num_workers: 16
|
||||||
|
|
||||||
|
train:
|
||||||
|
target: models.dataset.AnimeVideoDataset
|
||||||
|
params:
|
||||||
|
data_root: /data0/zhitong/datasets/animation_dataset
|
||||||
|
size: [320, 576]
|
||||||
|
motion_bucket_id: 160
|
||||||
|
fps_id: 6
|
||||||
|
num_frames: 15
|
||||||
|
cond_aug: False
|
||||||
|
nframe_range: [15, 200]
|
||||||
|
uncond_prob: 0.0
|
||||||
|
sketch_type: 'draw'
|
||||||
|
train_clips: 'train_clips_hist'
|
||||||
|
missing_controls: Null
|
||||||
|
sample_stride: 1
|
||||||
|
|
||||||
|
|
||||||
|
|
||||||
|
|
||||||
@@ -0,0 +1,197 @@
|
|||||||
|
import torch
|
||||||
|
from torch import nn
|
||||||
|
from einops import rearrange
|
||||||
|
import torch.nn.functional as F
|
||||||
|
|
||||||
|
from torch.backends.cuda import sdp_kernel
|
||||||
|
|
||||||
|
def remove_all_hooks(model: torch.nn.Module) -> None:
|
||||||
|
for child in model.children():
|
||||||
|
if hasattr(child, "_forward_hooks"):
|
||||||
|
child._forward_hooks.clear()
|
||||||
|
if hasattr(child, "_backward_hooks"):
|
||||||
|
child._backward_hooks.clear()
|
||||||
|
remove_all_hooks(child)
|
||||||
|
|
||||||
|
|
||||||
|
class Hacked_model(nn.Module):
|
||||||
|
def __init__(self, model, **kwargs):
|
||||||
|
super().__init__()
|
||||||
|
self.operator = Reference(model, **kwargs)
|
||||||
|
|
||||||
|
def forward(self, model, step, x_in, c_noise, cond_in, **additional_model_inputs):
|
||||||
|
# Register hooks
|
||||||
|
self.operator.register_hooks(model)
|
||||||
|
self.operator.setup(step)
|
||||||
|
# Model forward
|
||||||
|
out = model.apply_model(x_in, c_noise, cond_in, **additional_model_inputs)
|
||||||
|
# Remove hooks
|
||||||
|
self.operator.remove_hooks()
|
||||||
|
|
||||||
|
return out
|
||||||
|
|
||||||
|
def clear_storage(self):
|
||||||
|
self.operator.clear_storage()
|
||||||
|
|
||||||
|
|
||||||
|
class Operator():
|
||||||
|
def __init__(self, model):
|
||||||
|
self.hook_handles = []
|
||||||
|
self.layers = self.get_hook_layers(model)
|
||||||
|
|
||||||
|
def get_hook_layers(self, model):
|
||||||
|
raise NotImplementedError
|
||||||
|
|
||||||
|
def hook(self, module, inputs, outputs):
|
||||||
|
raise NotImplementedError
|
||||||
|
|
||||||
|
def setup(self, step, branch, opt):
|
||||||
|
raise NotImplementedError
|
||||||
|
|
||||||
|
def clear_storage(self):
|
||||||
|
self.storage.clear()
|
||||||
|
|
||||||
|
def register_hooks(self, model):
|
||||||
|
for m in model.modules():
|
||||||
|
index = id(m)
|
||||||
|
if index in self.layers.keys():
|
||||||
|
handle = m.register_forward_hook(self.hook)
|
||||||
|
self.hook_handles.append(handle)
|
||||||
|
|
||||||
|
def remove_hooks(self):
|
||||||
|
while len(self.hook_handles) > 0:
|
||||||
|
self.hook_handles[0].remove()
|
||||||
|
self.hook_handles.pop(0)
|
||||||
|
|
||||||
|
|
||||||
|
class Reference(Operator):
|
||||||
|
def __init__(self, model, **kwargs):
|
||||||
|
self.storage = nn.ParameterDict()
|
||||||
|
self.overlap = kwargs['overlap']
|
||||||
|
self.nframes = kwargs['nframes']
|
||||||
|
self.refattn_amp = kwargs['refattn_amplify']
|
||||||
|
self.refattn_hook = kwargs['refattn_hook']
|
||||||
|
self.prev_steps = kwargs['prev_steps']
|
||||||
|
super().__init__(model)
|
||||||
|
|
||||||
|
def setup(self, step):
|
||||||
|
self.step = step
|
||||||
|
|
||||||
|
def get_hook_layers(self, model):
|
||||||
|
layers = dict()
|
||||||
|
if self.refattn_hook:
|
||||||
|
# Hook ref attention layers in ControlNet
|
||||||
|
layer_name = model.control_model.spatial_self_attn_type.split('.')[-1]
|
||||||
|
i = 0
|
||||||
|
for name, module in model.control_model.named_modules():
|
||||||
|
if module.__class__.__name__ == layer_name and '.time_stack' not in name and '.attn1' in name:
|
||||||
|
layers[id(module)] = f'cnet-refcond-{i}'
|
||||||
|
i += 1
|
||||||
|
# Hook ref attention layers in UNet
|
||||||
|
layer_name = model.model.diffusion_model.spatial_self_attn_type.split('.')[-1]
|
||||||
|
i = 0
|
||||||
|
for name, module in model.model.diffusion_model.named_modules():
|
||||||
|
if module.__class__.__name__ == layer_name and '.time_stack' not in name and '.attn1' in name:
|
||||||
|
layers[id(module)] = f'unet-refcond-{i}'
|
||||||
|
i += 1
|
||||||
|
return layers
|
||||||
|
|
||||||
|
@torch.no_grad()
|
||||||
|
def hook(self, module, inputs, outputs):
|
||||||
|
layer = self.layers[id(module)]
|
||||||
|
if 'refcond' in layer:
|
||||||
|
out = self.reference_attn_forward(module, inputs, outputs)
|
||||||
|
return out
|
||||||
|
|
||||||
|
|
||||||
|
def reference_attn_forward(self, module, inputs, outputs):
|
||||||
|
overlap = self.overlap
|
||||||
|
T = self.nframes
|
||||||
|
h = module.heads
|
||||||
|
|
||||||
|
index = id(module)
|
||||||
|
layer = self.layers[index]
|
||||||
|
layer_ind = int(layer.split('-')[-1])
|
||||||
|
|
||||||
|
q = module.to_q(inputs[0])
|
||||||
|
k = module.to_k(inputs[0])
|
||||||
|
v = module.to_v(inputs[0])
|
||||||
|
|
||||||
|
olap = 3
|
||||||
|
|
||||||
|
if self.mode == 'normal':
|
||||||
|
indices = [
|
||||||
|
list(range(0, T)),
|
||||||
|
list(range(0, overlap+1)) + [0]*(T-overlap-1),
|
||||||
|
]
|
||||||
|
elif self.mode == 'prevref':
|
||||||
|
if self.step < self.prev_steps:
|
||||||
|
indices = [
|
||||||
|
list(range(0, 2*overlap+1)) + list(range(2*overlap+1-olap, T-olap)),
|
||||||
|
list(range(0, overlap+1)) + list(range(1, overlap+1)) + [0]*(T-2*overlap-1),
|
||||||
|
]
|
||||||
|
else:
|
||||||
|
'''indices = [
|
||||||
|
list(range(0, 2*overlap+1)) + list(range(2*overlap+1-olap, T-olap)),
|
||||||
|
list(range(0, overlap+1)) + [0]*overlap + [0]*(T-2*overlap-1),
|
||||||
|
]'''
|
||||||
|
indices = [
|
||||||
|
list(range(0, T)),
|
||||||
|
list(range(0, overlap+1)) + [0]*(T-overlap-1),
|
||||||
|
]
|
||||||
|
elif self.mode == 'normal1':
|
||||||
|
indices = [
|
||||||
|
list(range(0, T)),
|
||||||
|
list(range(0, overlap+1)) + [overlap] + [0]*(T-overlap-2),
|
||||||
|
]
|
||||||
|
'''elif self.mode == 'tempref':
|
||||||
|
if self.step < self.prev_steps:
|
||||||
|
indices = [
|
||||||
|
list(range(0, 2*overlap+1)) + [2*overlap]*olap + list(range(2*overlap+1, T-olap)),
|
||||||
|
list(range(0, overlap+1)) + list(range(1, overlap+1)) + [0]*(T-2*overlap-1),
|
||||||
|
]
|
||||||
|
else:
|
||||||
|
indices = [
|
||||||
|
list(range(0, 2*overlap+1)) + [2*overlap]*olap + list(range(2*overlap+1, T-olap)),
|
||||||
|
list(range(0, overlap+1)) + [0]*overlap + [0]*(T-2*overlap-1),
|
||||||
|
]'''
|
||||||
|
|
||||||
|
|
||||||
|
k = rearrange(k, '(b t) ... -> b t ...', t=T)
|
||||||
|
v = rearrange(v, '(b t) ... -> b t ...', t=T)
|
||||||
|
k = torch.cat([k[:, indices[i]] for i in range(len(indices))], dim=2).clone()
|
||||||
|
v = torch.cat([v[:, indices[i]] for i in range(len(indices))], dim=2).clone()
|
||||||
|
k = rearrange(k, 'b t ... -> (b t) ...')
|
||||||
|
v = rearrange(v, 'b t ... -> (b t) ...')
|
||||||
|
|
||||||
|
q, k, v = map(lambda t: rearrange(t, "b n (h d) -> b h n d", h=h), (q, k, v))
|
||||||
|
# Attention
|
||||||
|
N = q.shape[-2]
|
||||||
|
with sdp_kernel(**{"enable_math": True, "enable_flash": True, "enable_mem_efficient": True}):
|
||||||
|
if layer_ind > 12 or self.mode == 'normal':
|
||||||
|
attn_bias = None
|
||||||
|
else:
|
||||||
|
attn_bias = torch.zeros([T, 1, N, 2*N], device=q.device, dtype=torch.float32)
|
||||||
|
amplify = torch.tensor(self.refattn_amp).to(attn_bias)
|
||||||
|
amplify = rearrange(amplify, 'b t -> t 1 1 b')
|
||||||
|
amplify = amplify.log()
|
||||||
|
attn_bias[:, :, :, :N] = amplify[:, :, :, [0]]
|
||||||
|
attn_bias[:, :, :, N:] = amplify[:, :, :, [1]]
|
||||||
|
out = F.scaled_dot_product_attention(
|
||||||
|
q, k, v, attn_mask=attn_bias,
|
||||||
|
)
|
||||||
|
del q, k, v
|
||||||
|
|
||||||
|
out = rearrange(out, "b h n d -> b n (h d)", h=h)
|
||||||
|
return module.to_out(out)
|
||||||
|
|
||||||
|
|
||||||
|
|
||||||
|
|
||||||
|
|
||||||
|
|
||||||
|
|
||||||
|
|
||||||
|
|
||||||
|
|
||||||
|
|
||||||
@@ -0,0 +1,472 @@
|
|||||||
|
# @title Sampling function
|
||||||
|
import math
|
||||||
|
import os
|
||||||
|
from typing import Optional
|
||||||
|
import copy
|
||||||
|
|
||||||
|
import cv2
|
||||||
|
import numpy as np
|
||||||
|
import torch
|
||||||
|
from einops import rearrange, repeat
|
||||||
|
|
||||||
|
from ..sgm.util import append_dims
|
||||||
|
|
||||||
|
from .model_hack import Hacked_model, remove_all_hooks
|
||||||
|
from comfy.utils import ProgressBar
|
||||||
|
|
||||||
|
@torch.no_grad()
|
||||||
|
def sample_video(model, device, inp, arg, verbose=True):
|
||||||
|
|
||||||
|
def get_indices(n_samples, overlap):
|
||||||
|
indices = []
|
||||||
|
for n in range(n_samples):
|
||||||
|
if n == 0:
|
||||||
|
start = 1
|
||||||
|
first_ref = 0
|
||||||
|
second_refs = [0] * overlap
|
||||||
|
else:
|
||||||
|
start = end - overlap
|
||||||
|
first_ref = 0
|
||||||
|
second_refs = list(range(start, start+overlap))
|
||||||
|
|
||||||
|
end = start + arg.num_frames - overlap - 1
|
||||||
|
frame_ind = [first_ref] + second_refs + list(range(start, end))
|
||||||
|
|
||||||
|
ref_ind = 0
|
||||||
|
|
||||||
|
blend_ind = [0] + list(range(-overlap, 0))
|
||||||
|
|
||||||
|
indices.append([frame_ind, ref_ind, blend_ind])
|
||||||
|
|
||||||
|
return indices
|
||||||
|
|
||||||
|
remove_all_hooks(model)
|
||||||
|
overlap = arg.overlap
|
||||||
|
prev_attn_steps = arg.prev_attn_steps
|
||||||
|
|
||||||
|
n_samples = (len(inp.skts)-(arg.num_frames-overlap)) // (arg.num_frames-2*overlap-1) + 1
|
||||||
|
blend_indices = [ list(range(0, arg.num_frames)),
|
||||||
|
[0] + list(range(1, overlap+1))*2 + [overlap]*(arg.num_frames-2*overlap-1)]
|
||||||
|
blend_steps = [0]*(overlap+1) + [25]*(overlap) + [0]*(arg.num_frames-2*overlap-1)
|
||||||
|
indices = get_indices(n_samples=n_samples, overlap=overlap)
|
||||||
|
|
||||||
|
# Initialization
|
||||||
|
H, W = inp.imgs[0].shape[2:]
|
||||||
|
shape = (arg.num_frames, 4, H // 8, W // 8)
|
||||||
|
torch.manual_seed(arg.seed)
|
||||||
|
x_T = torch.randn(shape, dtype=torch.float32, device="cpu").to(device)
|
||||||
|
|
||||||
|
hacked = Hacked_model(
|
||||||
|
model, overlap=overlap, nframes=arg.num_frames,
|
||||||
|
refattn_hook=True, prev_steps=prev_attn_steps,
|
||||||
|
refattn_amplify = [
|
||||||
|
[1.0]*(overlap+1) + [1.0]*overlap + [1.0]*3 + [1.0]*7, # Self-attention
|
||||||
|
[1.0]*(overlap+1) + [10.0]*overlap + [1.0]*3 + [1.0]*7, # Ref-attention
|
||||||
|
]
|
||||||
|
)
|
||||||
|
|
||||||
|
first_cond = model.encode_first_stage(inp.imgs[0].to(device)) / model.scale_factor
|
||||||
|
first_conds = repeat(first_cond, 'b ... -> (b t) ...', t=arg.num_frames-overlap-1)
|
||||||
|
|
||||||
|
for i, index in enumerate(indices):
|
||||||
|
frame_ind, ref_ind, blend_ind = index
|
||||||
|
input_img = inp.imgs[ref_ind].to(device)
|
||||||
|
sketches = torch.cat([inp.skts[i] for i in frame_ind]).to(device)
|
||||||
|
if i == 0:
|
||||||
|
hacked.operator.mode = 'normal'
|
||||||
|
add_conds = None
|
||||||
|
intermediates = {'xt': None, 'denoised': None, 'x0': None}
|
||||||
|
else:
|
||||||
|
hacked.operator.mode = arg.ref_mode
|
||||||
|
prev_conds = x0[-overlap:] / model.scale_factor
|
||||||
|
add_conds = {'concat': {
|
||||||
|
'cond': torch.cat([ first_cond, prev_conds, first_conds ]),
|
||||||
|
} }
|
||||||
|
for k in intermediates['xt'].keys():
|
||||||
|
intermediates['denoised'][k] = intermediates['denoised'][k][blend_ind].clone()
|
||||||
|
|
||||||
|
x0, intermediates = sample(
|
||||||
|
model=model, device=device, x_T=x_T, input_img=input_img,
|
||||||
|
additional_conditions=add_conds, controls=sketches, hacked=hacked,
|
||||||
|
blend_x0=intermediates['denoised'], blend_ind=blend_indices, blend_steps=blend_steps,
|
||||||
|
return_intermediate=True, **vars(arg), verbose=True,
|
||||||
|
)
|
||||||
|
|
||||||
|
if i == 0:
|
||||||
|
outputs = torch.cat([first_cond*model.scale_factor, x0[-14:]]).cpu()
|
||||||
|
else:
|
||||||
|
outputs = torch.cat([outputs[:-overlap], x0[-14:].cpu()])
|
||||||
|
|
||||||
|
old_xT = x_T.clone()
|
||||||
|
x_T = torch.cat([ old_xT[[0]], old_xT[-overlap:], old_xT[-overlap:], old_xT[overlap+1:-overlap], ])
|
||||||
|
|
||||||
|
return outputs
|
||||||
|
|
||||||
|
|
||||||
|
@torch.no_grad()
|
||||||
|
def decode_video(model, device, latents, arg):
|
||||||
|
model.en_and_decode_n_samples_a_time = arg.decoding_t
|
||||||
|
|
||||||
|
N = latents.shape[0]
|
||||||
|
B = arg.decoding_t
|
||||||
|
olap = arg.decoding_olap
|
||||||
|
f = arg.decoding_first
|
||||||
|
|
||||||
|
end = 0
|
||||||
|
|
||||||
|
i = 0
|
||||||
|
|
||||||
|
with torch.autocast('cuda'):
|
||||||
|
while end < N:
|
||||||
|
start = i * (B - f - olap) + f
|
||||||
|
end = min( start + B - f, N)
|
||||||
|
|
||||||
|
indices = [0]*f + list(range(start, end))
|
||||||
|
|
||||||
|
inputs = latents[indices]
|
||||||
|
out = model.decode_first_stage(inputs.to(device)).cpu()
|
||||||
|
out = torch.clamp(out, min=-1.0, max=1.0)
|
||||||
|
if i == 0:
|
||||||
|
outputs = out.clone()
|
||||||
|
else:
|
||||||
|
outputs = torch.cat([ outputs, out[f+olap:] ])
|
||||||
|
i += 1
|
||||||
|
|
||||||
|
return outputs
|
||||||
|
|
||||||
|
|
||||||
|
|
||||||
|
|
||||||
|
def sample(
|
||||||
|
model,
|
||||||
|
device: str,
|
||||||
|
input_img: torch.Tensor,
|
||||||
|
hacked = None,
|
||||||
|
x_T: torch.Tensor = None,
|
||||||
|
num_frames: Optional[int] = None,
|
||||||
|
num_steps: Optional[int] = None,
|
||||||
|
palette: Optional[torch.Tensor] = None,
|
||||||
|
anchor: Optional[torch.Tensor] = None,
|
||||||
|
fps_id: int = 6,
|
||||||
|
motion_bucket_id: int = 127,
|
||||||
|
cond_aug: float = 0.02,
|
||||||
|
seed: int = 23,
|
||||||
|
decoding_t: int = 14, # Number of frames decoded at a time! This eats most VRAM. Reduce if necessary.
|
||||||
|
output_folder: Optional[str] = "/content/outputs",
|
||||||
|
verbose: bool = True,
|
||||||
|
controls: torch.Tensor = None,
|
||||||
|
blend_ind = None,
|
||||||
|
blend_x0: torch.Tensor = None,
|
||||||
|
scale = [1.0, 1.0],
|
||||||
|
return_intermediate: bool = False,
|
||||||
|
input_latent: torch.Tensor = None,
|
||||||
|
first_control: torch.Tensor = None,
|
||||||
|
blend_steps = None,
|
||||||
|
gamma = 0.0,
|
||||||
|
additional_conditions = None,
|
||||||
|
starting_conditions = None,
|
||||||
|
cfg_combine_forward = True,
|
||||||
|
**kwargs,
|
||||||
|
):
|
||||||
|
"""
|
||||||
|
Simple script to generate a single sample conditioned on an image `input_path` or multiple images, one for each
|
||||||
|
image file in folder `input_path`. If you run out of VRAM, try decreasing `decoding_t`.
|
||||||
|
"""
|
||||||
|
seed_everything(seed)
|
||||||
|
|
||||||
|
if True:
|
||||||
|
H, W = input_img.shape[2:]
|
||||||
|
assert input_img.shape[1] == 3
|
||||||
|
F = 8
|
||||||
|
C = 4
|
||||||
|
shape = (num_frames, C, H // F, W // F)
|
||||||
|
|
||||||
|
if motion_bucket_id > 255:
|
||||||
|
print("WARNING: High motion bucket! This may lead to suboptimal performance.")
|
||||||
|
if fps_id < 5:
|
||||||
|
print("WARNING: Small fps value! This may lead to suboptimal performance.")
|
||||||
|
if fps_id > 30:
|
||||||
|
print("WARNING: Large fps value! This may lead to suboptimal performance.")
|
||||||
|
|
||||||
|
value_dict = {}
|
||||||
|
value_dict["motion_bucket_id"] = motion_bucket_id
|
||||||
|
value_dict["fps_id"] = fps_id
|
||||||
|
value_dict["cond_aug"] = cond_aug
|
||||||
|
value_dict["cond_frames_without_noise"] = input_img
|
||||||
|
value_dict["cond_frames"] = input_img + cond_aug * torch.randn_like(input_img)
|
||||||
|
value_dict["cond_aug"] = cond_aug
|
||||||
|
model.sampler.verbose = verbose
|
||||||
|
model.sampler.device = device
|
||||||
|
|
||||||
|
with torch.no_grad():
|
||||||
|
with torch.autocast('cuda'):
|
||||||
|
# Prepare conditions
|
||||||
|
c, uc, additional_model_inputs = get_conditioning(
|
||||||
|
model,
|
||||||
|
get_unique_embedder_keys_from_conditioner(model.conditioner),
|
||||||
|
value_dict,
|
||||||
|
[1, num_frames],
|
||||||
|
T=num_frames,
|
||||||
|
input_latent=input_latent,
|
||||||
|
device=device,
|
||||||
|
controls=controls, palette=palette, anchor=anchor, first_control=first_control,
|
||||||
|
additional_conditions=additional_conditions,
|
||||||
|
)
|
||||||
|
# Initial noise
|
||||||
|
if x_T is None:
|
||||||
|
randn = torch.randn(shape, dtype=torch.float32, device="cpu").to(device)
|
||||||
|
else:
|
||||||
|
randn = x_T.clone()
|
||||||
|
|
||||||
|
# Prepare for swapping conditions
|
||||||
|
if starting_conditions is not None:
|
||||||
|
original_c = copy.deepcopy(c)
|
||||||
|
|
||||||
|
'''Sampling'''
|
||||||
|
intermediate = {'xt': {}, 'denoised': {},}
|
||||||
|
with torch.no_grad():
|
||||||
|
x = randn.clone()
|
||||||
|
sigmas = model.sampler.discretization(num_steps, device=device).to(torch.float32)
|
||||||
|
x *= torch.sqrt(1.0 + sigmas[0] ** 2.0)
|
||||||
|
num_sigmas = len(sigmas)
|
||||||
|
comfy_pbar = ProgressBar(num_sigmas)
|
||||||
|
for i in model.sampler.get_sigma_gen(num_sigmas):
|
||||||
|
|
||||||
|
# Blending
|
||||||
|
if blend_steps is not None and blend_ind is not None:
|
||||||
|
blend = (i < max(blend_steps))
|
||||||
|
target_ind = []
|
||||||
|
source_ind = []
|
||||||
|
for k, b in enumerate(blend_steps):
|
||||||
|
if i < b:
|
||||||
|
target_ind.append(blend_ind[0][k])
|
||||||
|
source_ind.append(blend_ind[1][k])
|
||||||
|
else:
|
||||||
|
blend = False
|
||||||
|
|
||||||
|
if return_intermediate:
|
||||||
|
intermediate['xt'][i] = x.clone()
|
||||||
|
|
||||||
|
if starting_conditions is not None:
|
||||||
|
if i < starting_conditions['step']:
|
||||||
|
c = copy.deepcopy(original_c)
|
||||||
|
for k in starting_conditions['cond'].keys():
|
||||||
|
c[k] = starting_conditions['cond'][k]
|
||||||
|
else:
|
||||||
|
c = original_c
|
||||||
|
|
||||||
|
if True:
|
||||||
|
# Prepare sigma
|
||||||
|
s_ones = x.new_ones([x.shape[0]], dtype=torch.float32)
|
||||||
|
sigma = s_ones * sigmas[i]
|
||||||
|
next_sigma = s_ones * sigmas[i+1]
|
||||||
|
sigma_hat = sigma * (gamma + 1.0)
|
||||||
|
# Denoising
|
||||||
|
denoised = denoise(
|
||||||
|
model, hacked, i, x, c, uc, additional_model_inputs,
|
||||||
|
sigma_hat, scale, cfg_combine_forward,
|
||||||
|
)
|
||||||
|
# CFG guidance
|
||||||
|
denoised = guidance(denoised, scale, num_frames)
|
||||||
|
if return_intermediate:
|
||||||
|
intermediate['denoised'][i] = denoised.clone()
|
||||||
|
|
||||||
|
# x0 blending
|
||||||
|
if blend and blend_x0 is not None:
|
||||||
|
#denoised[target_ind] = blend_x0[num_steps-1][source_ind]
|
||||||
|
denoised[target_ind] = blend_x0[i][source_ind]
|
||||||
|
|
||||||
|
# Euler step
|
||||||
|
d = (x - denoised) / append_dims(sigma_hat, x.ndim)
|
||||||
|
dt = append_dims(next_sigma - sigma_hat, x.ndim)
|
||||||
|
x = x + dt * d
|
||||||
|
comfy_pbar.update(1)
|
||||||
|
|
||||||
|
samples_z = x.clone().to(dtype=model.first_stage_model.dtype)
|
||||||
|
|
||||||
|
if return_intermediate:
|
||||||
|
return samples_z, intermediate
|
||||||
|
else:
|
||||||
|
return samples_z, None
|
||||||
|
|
||||||
|
|
||||||
|
def get_unique_embedder_keys_from_conditioner(conditioner):
|
||||||
|
return list(set([x.input_key for x in conditioner.embedders]))
|
||||||
|
|
||||||
|
def get_conditioning(model, keys, value_dict, N, T, device, input_latent, additional_conditions, dtype=None, **kwargs):
|
||||||
|
batch = {}
|
||||||
|
batch_uc = {}
|
||||||
|
|
||||||
|
for key in keys:
|
||||||
|
if key == "fps_id":
|
||||||
|
batch[key] = (
|
||||||
|
torch.tensor([value_dict["fps_id"]])
|
||||||
|
.to(device, dtype=dtype)
|
||||||
|
.repeat(int(math.prod(N)))
|
||||||
|
)
|
||||||
|
elif key == "motion_bucket_id":
|
||||||
|
batch[key] = (
|
||||||
|
torch.tensor([value_dict["motion_bucket_id"]])
|
||||||
|
.to(device, dtype=dtype)
|
||||||
|
.repeat(int(math.prod(N)))
|
||||||
|
)
|
||||||
|
elif key == "cond_aug":
|
||||||
|
batch[key] = repeat(
|
||||||
|
torch.tensor([value_dict["cond_aug"]]).to(device, dtype=dtype),
|
||||||
|
"1 -> b",
|
||||||
|
b=math.prod(N),
|
||||||
|
)
|
||||||
|
elif key == "cond_frames":
|
||||||
|
batch[key] = torch.cat([ value_dict["cond_frames"] ]*N[0])
|
||||||
|
elif key == "cond_frames_without_noise":
|
||||||
|
batch[key] = torch.cat([ value_dict["cond_frames_without_noise"] ]*N[0])
|
||||||
|
else:
|
||||||
|
batch[key] = value_dict[key]
|
||||||
|
|
||||||
|
if T is not None:
|
||||||
|
batch["num_video_frames"] = T
|
||||||
|
|
||||||
|
for key in batch.keys():
|
||||||
|
if key not in batch_uc and isinstance(batch[key], torch.Tensor):
|
||||||
|
batch_uc[key] = torch.clone(batch[key])
|
||||||
|
|
||||||
|
c, uc = model.conditioner.get_unconditional_conditioning(
|
||||||
|
batch,
|
||||||
|
batch_uc=batch_uc,
|
||||||
|
force_uc_zero_embeddings=[
|
||||||
|
"cond_frames",
|
||||||
|
"cond_frames_without_noise",
|
||||||
|
],
|
||||||
|
)
|
||||||
|
if input_latent is not None:
|
||||||
|
c['concat'] = input_latent.clone() / 0.18215
|
||||||
|
|
||||||
|
# from here, dtype is fp16
|
||||||
|
for k in ["crossattn", "concat"]:
|
||||||
|
uc[k] = repeat(uc[k], "b ... -> b t ...", t=T)
|
||||||
|
uc[k] = rearrange(uc[k], "b t ... -> (b t) ...", t=T)
|
||||||
|
c[k] = repeat(c[k], "b ... -> b t ...", t=T)
|
||||||
|
c[k] = rearrange(c[k], "b t ... -> (b t) ...", t=T)
|
||||||
|
for k in uc.keys():
|
||||||
|
uc[k] = uc[k].to(dtype=torch.float32)
|
||||||
|
c[k] = c[k].to(dtype=torch.float32)
|
||||||
|
|
||||||
|
if 'controls' in kwargs and kwargs['controls'] is not None:
|
||||||
|
uc['control_hint'] = kwargs['controls'].to(torch.float32)
|
||||||
|
c['control_hint'] = kwargs['controls'].to(torch.float32)
|
||||||
|
if 'first_control' in kwargs and kwargs['first_control'] is not None:
|
||||||
|
c['first_control'] = kwargs['first_control'].to(torch.float32)
|
||||||
|
uc['first_control'] = torch.zeros_like(c['first_control'])
|
||||||
|
if 'palette' in kwargs and kwargs['palette'] is not None:
|
||||||
|
uc['palette'] = kwargs['palette'].to(torch.float32)
|
||||||
|
c['palette'] = kwargs['palette'].to(torch.float32)
|
||||||
|
if 'anchor' in kwargs and kwargs['anchor'] is not None:
|
||||||
|
uc['anchor'] = kwargs['anchor'].to(torch.float32)
|
||||||
|
c['anchor'] = kwargs['anchor'].to(torch.float32)
|
||||||
|
|
||||||
|
if additional_conditions is not None:
|
||||||
|
for k in additional_conditions.keys():
|
||||||
|
c[k] = additional_conditions[k]['cond'].to(torch.float32)
|
||||||
|
if 'uncond' in additional_conditions[k].keys():
|
||||||
|
uc[k] = additional_conditions[k]['uncond'].to(torch.float32)
|
||||||
|
else:
|
||||||
|
uc[k] = additional_conditions[k]['cond'].to(torch.float32)
|
||||||
|
|
||||||
|
additional_model_inputs = {}
|
||||||
|
additional_model_inputs["image_only_indicator"] = torch.zeros(1, T).to(device)
|
||||||
|
additional_model_inputs["num_video_frames"] = batch["num_video_frames"]
|
||||||
|
|
||||||
|
for k in additional_model_inputs:
|
||||||
|
if isinstance(additional_model_inputs[k], torch.Tensor):
|
||||||
|
additional_model_inputs[k] = additional_model_inputs[k].to(dtype=torch.float32)
|
||||||
|
|
||||||
|
return c, uc, additional_model_inputs
|
||||||
|
|
||||||
|
def denoise(
|
||||||
|
model, hacked, step, x,
|
||||||
|
c, uc, additional_model_inputs,
|
||||||
|
sigma_hat, scale, cfg_combine_forward,
|
||||||
|
):
|
||||||
|
# Prepare model input
|
||||||
|
if scale[1] != 1.0 and cfg_combine_forward:
|
||||||
|
cond_in = dict()
|
||||||
|
if additional_model_inputs['image_only_indicator'].shape[0] == 1:
|
||||||
|
additional_model_inputs["image_only_indicator"] = additional_model_inputs["image_only_indicator"].repeat(2, 1)
|
||||||
|
for k in c:
|
||||||
|
if k in ["vector", "crossattn", "concat"] + model.sampler.guider.additional_cond_keys:
|
||||||
|
cond_in[k] = torch.cat((uc[k], c[k]), 0)
|
||||||
|
else:
|
||||||
|
assert c[k] == uc[k]
|
||||||
|
cond_in[k] = c[k]
|
||||||
|
x_in = torch.cat([x] * 2)
|
||||||
|
s_in = torch.cat([sigma_hat] * 2)
|
||||||
|
else:
|
||||||
|
cond_in = c
|
||||||
|
x_in = x
|
||||||
|
s_in = sigma_hat
|
||||||
|
|
||||||
|
if hacked is not None:
|
||||||
|
model_forward = lambda inp, c_noise, cond, **add: hacked(model, step, inp, c_noise, cond, **add)
|
||||||
|
else:
|
||||||
|
model_forward = model.apply_model
|
||||||
|
|
||||||
|
denoised = model.denoiser(model_forward, x_in, s_in, cond_in, **additional_model_inputs)
|
||||||
|
|
||||||
|
if not cfg_combine_forward and scale[1] != 1.0:
|
||||||
|
uc_denoised = model.denoiser(model_forward, x_in, s_in, uc, **additional_model_inputs)
|
||||||
|
denoised = torch.cat([uc_denoised, denoised])
|
||||||
|
|
||||||
|
if denoised.shape[0] < x_in.shape[0]:
|
||||||
|
denoised = rearrange(denoised, '(b t) ... -> b t ...', t=additional_model_inputs["num_video_frames"]-1)
|
||||||
|
denoised = torch.cat([denoised[:, [0]], denoised], dim=1)
|
||||||
|
denoised = rearrange(denoised, 'b t ... -> (b t) ...')
|
||||||
|
|
||||||
|
return denoised
|
||||||
|
|
||||||
|
def guidance(denoised, scale, num_frames):
|
||||||
|
if scale[1] != 1.0:
|
||||||
|
x_u, x_c = denoised.chunk(2)
|
||||||
|
x_u = rearrange(x_u, "(b t) ... -> b t ...", t=num_frames)
|
||||||
|
x_c = rearrange(x_c, "(b t) ... -> b t ...", t=num_frames)
|
||||||
|
scales = torch.linspace(scale[0], scale[1], num_frames).unsqueeze(0)
|
||||||
|
scales = repeat(scales, "1 t -> b t", b=x_u.shape[0])
|
||||||
|
scales = append_dims(scales, x_u.ndim).to(x_u.device)
|
||||||
|
denoised = rearrange(x_u + scales * (x_c - x_u), "b t ... -> (b t) ...")
|
||||||
|
|
||||||
|
return denoised
|
||||||
|
|
||||||
|
|
||||||
|
def write_video(output_folder, fps_id, samples):
|
||||||
|
os.makedirs(output_folder, exist_ok=True)
|
||||||
|
video_path = os.path.join(output_folder, f".mp4")
|
||||||
|
writer = cv2.VideoWriter(
|
||||||
|
video_path,
|
||||||
|
cv2.VideoWriter_fourcc(*"MP4V"),
|
||||||
|
fps_id + 1,
|
||||||
|
(samples.shape[-1], samples.shape[-2]),
|
||||||
|
)
|
||||||
|
|
||||||
|
vid = (
|
||||||
|
(rearrange(samples, "t c h w -> t h w c") * 255)
|
||||||
|
.cpu()
|
||||||
|
.numpy()
|
||||||
|
.astype(np.uint8)
|
||||||
|
)
|
||||||
|
for frame in vid:
|
||||||
|
frame = cv2.cvtColor(frame, cv2.COLOR_RGB2BGR)
|
||||||
|
writer.write(frame)
|
||||||
|
writer.release()
|
||||||
|
|
||||||
|
|
||||||
|
def seed_everything(seed: int):
|
||||||
|
import random, os
|
||||||
|
import numpy as np
|
||||||
|
import torch
|
||||||
|
|
||||||
|
random.seed(seed)
|
||||||
|
os.environ['PYTHONHASHSEED'] = str(seed)
|
||||||
|
np.random.seed(seed)
|
||||||
|
torch.manual_seed(seed)
|
||||||
|
torch.cuda.manual_seed(seed)
|
||||||
|
torch.backends.cudnn.deterministic = True
|
||||||
|
torch.backends.cudnn.benchmark = True
|
||||||
|
After Width: | Height: | Size: 103 KiB |
|
After Width: | Height: | Size: 103 KiB |
|
After Width: | Height: | Size: 100 KiB |
|
After Width: | Height: | Size: 98 KiB |
|
After Width: | Height: | Size: 98 KiB |
|
After Width: | Height: | Size: 99 KiB |
|
After Width: | Height: | Size: 99 KiB |
|
After Width: | Height: | Size: 101 KiB |
|
After Width: | Height: | Size: 101 KiB |
|
After Width: | Height: | Size: 102 KiB |
|
After Width: | Height: | Size: 101 KiB |
|
After Width: | Height: | Size: 103 KiB |
|
After Width: | Height: | Size: 104 KiB |
|
After Width: | Height: | Size: 104 KiB |
|
After Width: | Height: | Size: 100 KiB |
|
After Width: | Height: | Size: 99 KiB |
|
After Width: | Height: | Size: 98 KiB |
|
After Width: | Height: | Size: 98 KiB |
|
After Width: | Height: | Size: 98 KiB |
|
After Width: | Height: | Size: 98 KiB |
|
After Width: | Height: | Size: 98 KiB |
|
After Width: | Height: | Size: 98 KiB |
|
After Width: | Height: | Size: 99 KiB |
|
After Width: | Height: | Size: 99 KiB |
|
After Width: | Height: | Size: 99 KiB |
|
After Width: | Height: | Size: 99 KiB |
|
After Width: | Height: | Size: 100 KiB |
|
After Width: | Height: | Size: 99 KiB |
|
After Width: | Height: | Size: 99 KiB |
|
After Width: | Height: | Size: 102 KiB |
|
After Width: | Height: | Size: 102 KiB |
|
After Width: | Height: | Size: 105 KiB |
|
After Width: | Height: | Size: 107 KiB |
|
After Width: | Height: | Size: 106 KiB |
|
After Width: | Height: | Size: 105 KiB |
|
After Width: | Height: | Size: 104 KiB |
|
After Width: | Height: | Size: 105 KiB |
|
After Width: | Height: | Size: 106 KiB |
|
After Width: | Height: | Size: 106 KiB |
|
After Width: | Height: | Size: 106 KiB |
|
After Width: | Height: | Size: 104 KiB |
|
After Width: | Height: | Size: 104 KiB |
|
After Width: | Height: | Size: 104 KiB |
|
After Width: | Height: | Size: 104 KiB |
|
After Width: | Height: | Size: 108 KiB |
|
After Width: | Height: | Size: 107 KiB |
|
After Width: | Height: | Size: 108 KiB |
|
After Width: | Height: | Size: 109 KiB |
|
After Width: | Height: | Size: 111 KiB |
|
After Width: | Height: | Size: 107 KiB |
|
After Width: | Height: | Size: 106 KiB |
|
After Width: | Height: | Size: 107 KiB |
|
After Width: | Height: | Size: 107 KiB |
|
After Width: | Height: | Size: 107 KiB |
|
After Width: | Height: | Size: 107 KiB |
|
After Width: | Height: | Size: 107 KiB |
|
After Width: | Height: | Size: 107 KiB |
|
After Width: | Height: | Size: 107 KiB |
|
After Width: | Height: | Size: 107 KiB |
|
After Width: | Height: | Size: 108 KiB |
|
After Width: | Height: | Size: 107 KiB |
|
After Width: | Height: | Size: 107 KiB |
|
After Width: | Height: | Size: 106 KiB |
|
After Width: | Height: | Size: 106 KiB |
|
After Width: | Height: | Size: 105 KiB |
|
After Width: | Height: | Size: 105 KiB |
|
After Width: | Height: | Size: 105 KiB |
|
After Width: | Height: | Size: 106 KiB |
|
After Width: | Height: | Size: 106 KiB |
|
After Width: | Height: | Size: 106 KiB |
|
After Width: | Height: | Size: 106 KiB |
|
After Width: | Height: | Size: 106 KiB |
|
After Width: | Height: | Size: 107 KiB |
|
After Width: | Height: | Size: 107 KiB |
|
After Width: | Height: | Size: 107 KiB |
@@ -0,0 +1,301 @@
|
|||||||
|
import logging
|
||||||
|
from inspect import isfunction
|
||||||
|
|
||||||
|
import torch
|
||||||
|
import torch.nn.functional as F
|
||||||
|
from einops import rearrange, repeat
|
||||||
|
from packaging import version
|
||||||
|
from torch import nn
|
||||||
|
|
||||||
|
logpy = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
if version.parse(torch.__version__) >= version.parse("2.0.0"):
|
||||||
|
SDP_IS_AVAILABLE = True
|
||||||
|
from torch.backends.cuda import SDPBackend, sdp_kernel
|
||||||
|
|
||||||
|
BACKEND_MAP = {
|
||||||
|
SDPBackend.MATH: {
|
||||||
|
"enable_math": True,
|
||||||
|
"enable_flash": False,
|
||||||
|
"enable_mem_efficient": False,
|
||||||
|
},
|
||||||
|
SDPBackend.FLASH_ATTENTION: {
|
||||||
|
"enable_math": False,
|
||||||
|
"enable_flash": True,
|
||||||
|
"enable_mem_efficient": False,
|
||||||
|
},
|
||||||
|
SDPBackend.EFFICIENT_ATTENTION: {
|
||||||
|
"enable_math": False,
|
||||||
|
"enable_flash": False,
|
||||||
|
"enable_mem_efficient": True,
|
||||||
|
},
|
||||||
|
None: {"enable_math": True, "enable_flash": True, "enable_mem_efficient": True},
|
||||||
|
}
|
||||||
|
else:
|
||||||
|
from contextlib import nullcontext
|
||||||
|
|
||||||
|
SDP_IS_AVAILABLE = False
|
||||||
|
sdp_kernel = nullcontext
|
||||||
|
BACKEND_MAP = {}
|
||||||
|
logpy.warn(
|
||||||
|
f"No SDP backend available, likely because you are running in pytorch "
|
||||||
|
f"versions < 2.0. In fact, you are using PyTorch {torch.__version__}. "
|
||||||
|
f"You might want to consider upgrading."
|
||||||
|
)
|
||||||
|
|
||||||
|
try:
|
||||||
|
import xformers
|
||||||
|
import xformers.ops
|
||||||
|
|
||||||
|
XFORMERS_IS_AVAILABLE = True
|
||||||
|
except:
|
||||||
|
XFORMERS_IS_AVAILABLE = False
|
||||||
|
logpy.warn("no module 'xformers'. Processing without...")
|
||||||
|
|
||||||
|
|
||||||
|
'''This temporal attention replace the original one in SVD to disable the temporal
|
||||||
|
attentions between the first frame (reference path) and the remaining 14 frames (video path).'''
|
||||||
|
class TemporalAttention_Masked(nn.Module):
|
||||||
|
def __init__(
|
||||||
|
self,
|
||||||
|
query_dim,
|
||||||
|
context_dim=None,
|
||||||
|
heads=8,
|
||||||
|
dim_head=64,
|
||||||
|
dropout=0.0,
|
||||||
|
backend=None,
|
||||||
|
):
|
||||||
|
super().__init__()
|
||||||
|
inner_dim = dim_head * heads
|
||||||
|
context_dim = default(context_dim, query_dim)
|
||||||
|
|
||||||
|
self.scale = dim_head**-0.5
|
||||||
|
self.heads = heads
|
||||||
|
|
||||||
|
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_out = nn.Sequential(
|
||||||
|
nn.Linear(inner_dim, query_dim), nn.Dropout(dropout)
|
||||||
|
)
|
||||||
|
self.backend = backend
|
||||||
|
|
||||||
|
def forward(
|
||||||
|
self,
|
||||||
|
x,
|
||||||
|
context=None,
|
||||||
|
mask=None,
|
||||||
|
additional_tokens=None,
|
||||||
|
n_times_crossframe_attn_in_self=0,
|
||||||
|
):
|
||||||
|
if hasattr(self, '_forward_hooks') and len(self._forward_hooks) > 0:
|
||||||
|
# If hooked do nothing
|
||||||
|
return x
|
||||||
|
else:
|
||||||
|
return self._forward(x, context, mask, additional_tokens, n_times_crossframe_attn_in_self)
|
||||||
|
|
||||||
|
def _forward(
|
||||||
|
self,
|
||||||
|
x,
|
||||||
|
context=None,
|
||||||
|
mask=None,
|
||||||
|
additional_tokens=None,
|
||||||
|
n_times_crossframe_attn_in_self=0,
|
||||||
|
):
|
||||||
|
h = self.heads
|
||||||
|
|
||||||
|
if mask is None:
|
||||||
|
T = x.shape[-2]
|
||||||
|
dt = T - 14
|
||||||
|
mask = torch.ones(T, T).to(x)
|
||||||
|
mask[:, :dt] = 0.0
|
||||||
|
mask[:dt, :] = 0.0
|
||||||
|
inds = [t for t in range(dt)]
|
||||||
|
mask[inds, inds] = 1.0
|
||||||
|
mask = rearrange(mask, 'h w -> 1 1 h w')
|
||||||
|
mask = mask.bool()
|
||||||
|
|
||||||
|
if additional_tokens is not None:
|
||||||
|
# get the number of masked tokens at the beginning of the output sequence
|
||||||
|
n_tokens_to_mask = additional_tokens.shape[1]
|
||||||
|
# add additional token
|
||||||
|
x = torch.cat([additional_tokens, x], dim=1)
|
||||||
|
|
||||||
|
q = self.to_q(x)
|
||||||
|
context = default(context, x)
|
||||||
|
k = self.to_k(context)
|
||||||
|
v = self.to_v(context)
|
||||||
|
|
||||||
|
if n_times_crossframe_attn_in_self:
|
||||||
|
# reprogramming cross-frame attention as in https://arxiv.org/abs/2303.13439
|
||||||
|
assert x.shape[0] % n_times_crossframe_attn_in_self == 0
|
||||||
|
n_cp = x.shape[0] // n_times_crossframe_attn_in_self
|
||||||
|
k = repeat(
|
||||||
|
k[::n_times_crossframe_attn_in_self], "b ... -> (b n) ...", n=n_cp
|
||||||
|
)
|
||||||
|
v = repeat(
|
||||||
|
v[::n_times_crossframe_attn_in_self], "b ... -> (b n) ...", n=n_cp
|
||||||
|
)
|
||||||
|
|
||||||
|
q, k, v = map(lambda t: rearrange(t, "b n (h d) -> b h n d", h=h), (q, k, v))
|
||||||
|
|
||||||
|
with sdp_kernel(**BACKEND_MAP[self.backend]):
|
||||||
|
# print("dispatching into backend", self.backend, "q/k/v shape: ", q.shape, k.shape, v.shape)
|
||||||
|
out = F.scaled_dot_product_attention(
|
||||||
|
q, k, v, attn_mask=mask
|
||||||
|
) # scale is dim_head ** -0.5 per default
|
||||||
|
|
||||||
|
del q, k, v
|
||||||
|
out = rearrange(out, "b h n d -> b n (h d)", h=h)
|
||||||
|
|
||||||
|
if additional_tokens is not None:
|
||||||
|
# remove additional token
|
||||||
|
out = out[:, n_tokens_to_mask:]
|
||||||
|
return self.to_out(out)
|
||||||
|
|
||||||
|
|
||||||
|
'''The reference attention which replace the original spatial self-attention layers in SVD.'''
|
||||||
|
class ReferenceAttention(nn.Module):
|
||||||
|
def __init__(
|
||||||
|
self,
|
||||||
|
query_dim,
|
||||||
|
context_dim=None,
|
||||||
|
heads=8,
|
||||||
|
dim_head=64,
|
||||||
|
dropout=0.0,
|
||||||
|
backend=None,
|
||||||
|
):
|
||||||
|
super().__init__()
|
||||||
|
inner_dim = dim_head * heads
|
||||||
|
context_dim = default(context_dim, query_dim)
|
||||||
|
|
||||||
|
self.scale = dim_head**-0.5
|
||||||
|
self.heads = heads
|
||||||
|
|
||||||
|
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_out = nn.Sequential(
|
||||||
|
nn.Linear(inner_dim, query_dim), nn.Dropout(dropout)
|
||||||
|
)
|
||||||
|
self.backend = backend
|
||||||
|
|
||||||
|
def forward(
|
||||||
|
self,
|
||||||
|
x,
|
||||||
|
context=None,
|
||||||
|
mask=None,
|
||||||
|
additional_tokens=None,
|
||||||
|
n_times_crossframe_attn_in_self=0,
|
||||||
|
):
|
||||||
|
if hasattr(self, '_forward_hooks') and len(self._forward_hooks) > 0:
|
||||||
|
# If hooked do nothing
|
||||||
|
return x
|
||||||
|
else:
|
||||||
|
return self._forward(x, context, mask, additional_tokens, n_times_crossframe_attn_in_self)
|
||||||
|
|
||||||
|
def _forward(
|
||||||
|
self,
|
||||||
|
x,
|
||||||
|
context=None,
|
||||||
|
mask=None,
|
||||||
|
additional_tokens=None,
|
||||||
|
n_times_crossframe_attn_in_self=0,
|
||||||
|
):
|
||||||
|
B = x.shape[0] // 14
|
||||||
|
T = x.shape[0] // B
|
||||||
|
h = self.heads
|
||||||
|
|
||||||
|
if additional_tokens is not None:
|
||||||
|
# get the number of masked tokens at the beginning of the output sequence
|
||||||
|
n_tokens_to_mask = additional_tokens.shape[1]
|
||||||
|
# add additional token
|
||||||
|
x = torch.cat([additional_tokens, x], dim=1)
|
||||||
|
|
||||||
|
q = self.to_q(x)
|
||||||
|
context = default(context, x)
|
||||||
|
k = self.to_k(context)
|
||||||
|
v = self.to_v(context)
|
||||||
|
# Refconcat: Q [K, K0] [V, V0]
|
||||||
|
k0 = rearrange(k, '(b t) ... -> b t ...', t=T)[:, [0]]
|
||||||
|
k0 = repeat(k0, 'b t0 ... -> b (t t0) ...', t=T)
|
||||||
|
k0 = rearrange(k0, 'b t ... -> (b t) ...')
|
||||||
|
v0 = rearrange(v, '(b t) ... -> b t ...', t=T)[:, [0]]
|
||||||
|
v0 = repeat(v0, 'b t0 ... -> b (t t0) ...', t=T)
|
||||||
|
v0 = rearrange(v0, 'b t ... -> (b t) ...')
|
||||||
|
k = torch.cat([k, k0], dim=1)
|
||||||
|
v = torch.cat([v, v0], dim=1)
|
||||||
|
|
||||||
|
if n_times_crossframe_attn_in_self:
|
||||||
|
# reprogramming cross-frame attention as in https://arxiv.org/abs/2303.13439
|
||||||
|
assert x.shape[0] % n_times_crossframe_attn_in_self == 0
|
||||||
|
n_cp = x.shape[0] // n_times_crossframe_attn_in_self
|
||||||
|
k = repeat(
|
||||||
|
k[::n_times_crossframe_attn_in_self], "b ... -> (b n) ...", n=n_cp
|
||||||
|
)
|
||||||
|
v = repeat(
|
||||||
|
v[::n_times_crossframe_attn_in_self], "b ... -> (b n) ...", n=n_cp
|
||||||
|
)
|
||||||
|
|
||||||
|
q, k, v = map(lambda t: rearrange(t, "b n (h d) -> b h n d", h=h), (q, k, v))
|
||||||
|
|
||||||
|
with sdp_kernel(**BACKEND_MAP[self.backend]):
|
||||||
|
# print("dispatching into backend", self.backend, "q/k/v shape: ", q.shape, k.shape, v.shape)
|
||||||
|
out = F.scaled_dot_product_attention(
|
||||||
|
q, k, v, attn_mask=mask
|
||||||
|
) # scale is dim_head ** -0.5 per default
|
||||||
|
|
||||||
|
del q, k, v
|
||||||
|
out = rearrange(out, "b h n d -> b n (h d)", h=h)
|
||||||
|
|
||||||
|
if additional_tokens is not None:
|
||||||
|
# remove additional token
|
||||||
|
out = out[:, n_tokens_to_mask:]
|
||||||
|
return self.to_out(out)
|
||||||
|
|
||||||
|
|
||||||
|
'''The 3D convolutional layers which disables the interactions between the
|
||||||
|
first frame (reference path) and the remaining 14 frames (video path).'''
|
||||||
|
class Conv3d_Masked(nn.Module):
|
||||||
|
def __init__(self, in_channels, out_channels, kernel_size, padding):
|
||||||
|
super().__init__()
|
||||||
|
self.padding = padding
|
||||||
|
self.weight = nn.Parameter( torch.zeros([out_channels, in_channels, kernel_size[0], kernel_size[1], kernel_size[2]]) )
|
||||||
|
self.bias = nn.Parameter( torch.zeros([out_channels]) )
|
||||||
|
|
||||||
|
def forward(self, x):
|
||||||
|
dt = x.shape[2] - 14
|
||||||
|
|
||||||
|
zeros_pad = torch.zeros_like(x[:, :, [0]])
|
||||||
|
|
||||||
|
xs = []
|
||||||
|
for i in range(dt):
|
||||||
|
xs.append( x[:, :, [i]] )
|
||||||
|
xs.append( zeros_pad )
|
||||||
|
xs.append( x[:, :, dt:] )
|
||||||
|
x = torch.cat(xs, dim=2)
|
||||||
|
|
||||||
|
x = torch.nn.functional.conv3d(
|
||||||
|
input=x,
|
||||||
|
weight=self.weight,
|
||||||
|
bias=self.bias,
|
||||||
|
padding=self.padding,
|
||||||
|
)
|
||||||
|
|
||||||
|
out_ind = [2*i for i in range(dt)]
|
||||||
|
|
||||||
|
x = torch.cat([ x[:, :, out_ind], x[:, :, 2*dt:] ], dim=2)
|
||||||
|
|
||||||
|
return x
|
||||||
|
|
||||||
|
|
||||||
|
|
||||||
|
def default(val, d):
|
||||||
|
if exists(val):
|
||||||
|
return val
|
||||||
|
return d() if isfunction(d) else d
|
||||||
|
|
||||||
|
def exists(val):
|
||||||
|
return val is not None
|
||||||
@@ -0,0 +1,209 @@
|
|||||||
|
import os
|
||||||
|
import torch
|
||||||
|
import folder_paths
|
||||||
|
import comfy.model_management as mm
|
||||||
|
|
||||||
|
import argparse
|
||||||
|
from omegaconf import OmegaConf
|
||||||
|
import logging
|
||||||
|
from .sgm.util import instantiate_from_config
|
||||||
|
from .inference.sample_func import sample_video, decode_video
|
||||||
|
|
||||||
|
logging.basicConfig(level=logging.INFO, format='%(asctime)s - %(levelname)s - %(message)s')
|
||||||
|
log = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
script_directory = os.path.dirname(os.path.abspath(__file__))
|
||||||
|
|
||||||
|
class LoadLVCDModel:
|
||||||
|
@classmethod
|
||||||
|
def INPUT_TYPES(s):
|
||||||
|
return {
|
||||||
|
"required": {
|
||||||
|
"model": (folder_paths.get_filename_list("checkpoints"), {"tooltip": "Normal SVD model, default is the normal very first non-XT SVD"} ),
|
||||||
|
"use_xformers": ("BOOLEAN", {"default": False}),
|
||||||
|
},
|
||||||
|
"optional": {
|
||||||
|
"precision": (["fp16", "fp32", "bf16"],
|
||||||
|
{"default": "fp16"}
|
||||||
|
),
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
RETURN_TYPES = ("LVCDPIPE",)
|
||||||
|
RETURN_NAMES = ("LVCD_pipe", )
|
||||||
|
FUNCTION = "loadmodel"
|
||||||
|
CATEGORY = "ComfyUI-LVCDWrapper"
|
||||||
|
|
||||||
|
def loadmodel(self, model, precision, use_xformers):
|
||||||
|
|
||||||
|
device = mm.get_torch_device()
|
||||||
|
print(device)
|
||||||
|
offload_device = mm.unet_offload_device()
|
||||||
|
dtype = {"bf16": torch.bfloat16, "fp16": torch.float16, "fp32": torch.float32}[precision]
|
||||||
|
mm.soft_empty_cache()
|
||||||
|
|
||||||
|
svd_model_path = folder_paths.get_full_path_or_raise("checkpoints", model)
|
||||||
|
download_path = os.path.join(folder_paths.models_dir, "lvcd")
|
||||||
|
lvcd_path = os.path.join(download_path, "lvcd-fp16.safetensors")
|
||||||
|
|
||||||
|
if not os.path.exists(lvcd_path):
|
||||||
|
log.info(f"Downloading LVCD model to: {lvcd_path}")
|
||||||
|
from huggingface_hub import snapshot_download
|
||||||
|
|
||||||
|
snapshot_download(
|
||||||
|
repo_id="Kijai/LVCD-pruned",
|
||||||
|
local_dir=download_path,
|
||||||
|
local_dir_use_symlinks=False,
|
||||||
|
)
|
||||||
|
|
||||||
|
config_path = os.path.join(script_directory, "configs", "lvcd.yaml")
|
||||||
|
config = OmegaConf.load(config_path)
|
||||||
|
config.model.params.drop_first_stage_model = False
|
||||||
|
config.model.params.init_from_unet = False
|
||||||
|
print(config.model.params.conditioner_config.params.emb_models[0])
|
||||||
|
|
||||||
|
if use_xformers:
|
||||||
|
config.model.params.network_config.params.spatial_transformer_attn_type = 'softmax-xformers'
|
||||||
|
config.model.params.controlnet_config.params.spatial_transformer_attn_type = 'softmax-xformers'
|
||||||
|
config.model.params.conditioner_config.params.emb_models[3].params.encoder_config.params.ddconfig.attn_type = 'vanilla-xformers'
|
||||||
|
else:
|
||||||
|
config.model.params.network_config.params.spatial_transformer_attn_type = 'softmax'
|
||||||
|
config.model.params.controlnet_config.params.spatial_transformer_attn_type = 'softmax'
|
||||||
|
config.model.params.conditioner_config.params.emb_models[3].params.encoder_config.params.ddconfig.attn_type = 'vanilla'
|
||||||
|
|
||||||
|
config.model.params.ckpt_path = svd_model_path
|
||||||
|
config.model.params.control_model_path = lvcd_path
|
||||||
|
|
||||||
|
with torch.device(device):
|
||||||
|
model = instantiate_from_config(config.model).to(device).eval().requires_grad_(False)
|
||||||
|
|
||||||
|
model.model.to(dtype)
|
||||||
|
model.control_model.to(dtype)
|
||||||
|
model.eval()
|
||||||
|
model = model.requires_grad_(False)
|
||||||
|
|
||||||
|
lvcd_pipe = {
|
||||||
|
"model": model,
|
||||||
|
"dtype": dtype,
|
||||||
|
}
|
||||||
|
|
||||||
|
return (lvcd_pipe,)
|
||||||
|
|
||||||
|
class LVCDSampler:
|
||||||
|
@classmethod
|
||||||
|
def INPUT_TYPES(s):
|
||||||
|
return {
|
||||||
|
"required": {
|
||||||
|
"LVCD_pipe": ("LVCDPIPE",),
|
||||||
|
"ref_images": ("IMAGE",),
|
||||||
|
"sketch_images": ("IMAGE",),
|
||||||
|
"num_frames": ("INT", {"default": 19, "min": 1, "max": 100, "step": 1}),
|
||||||
|
"num_steps": ("INT", {"default": 25, "min": 1, "max": 100, "step": 1}),
|
||||||
|
"fps_id": ("INT", {"default": 6, "min": 1, "max": 100, "step": 1}),
|
||||||
|
"motion_bucket_id": ("INT", {"default": 160, "min": 0, "max": 1000, "step": 1}),
|
||||||
|
"cond_aug": ("FLOAT", {"default": 0.0, "min": 0.0, "max": 1.0, "step": 0.01}),
|
||||||
|
"overlap": ("INT", {"default": 4, "min": 1, "max": 100, "step": 1}),
|
||||||
|
"prev_attn_steps": ("INT", {"default": 25, "min": 1, "max": 100, "step": 1}),
|
||||||
|
"seed": ("INT", {"default": 123, "min": 0, "max": 2**32, "step": 1}),
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
RETURN_TYPES = ("LVCDPIPE", "SVDSAMPLES",)
|
||||||
|
RETURN_NAMES = ("LVCD_pipe", "samples",)
|
||||||
|
FUNCTION = "loadmodel"
|
||||||
|
CATEGORY = "ComfyUI-LVCDWrapper"
|
||||||
|
|
||||||
|
def loadmodel(self, LVCD_pipe, ref_images, sketch_images, num_frames, num_steps, fps_id, motion_bucket_id, cond_aug, overlap,
|
||||||
|
prev_attn_steps, seed):
|
||||||
|
|
||||||
|
device = mm.get_torch_device()
|
||||||
|
offload_device = mm.unet_offload_device()
|
||||||
|
mm.soft_empty_cache()
|
||||||
|
|
||||||
|
model = LVCD_pipe["model"]
|
||||||
|
|
||||||
|
inp = argparse.ArgumentParser()
|
||||||
|
B, H, W, C = ref_images.shape
|
||||||
|
inp.resolution = [H, W]
|
||||||
|
|
||||||
|
inp.imgs = []
|
||||||
|
inp.skts = []
|
||||||
|
|
||||||
|
ref_images = ref_images.permute(0, 3, 1, 2).to(device) * 2 - 1
|
||||||
|
for ref_img in ref_images:
|
||||||
|
print(ref_img.shape)
|
||||||
|
inp.imgs.append(ref_img.unsqueeze(0))
|
||||||
|
sketch_images = sketch_images.permute(0, 3, 1, 2).to(device)
|
||||||
|
for skt in sketch_images:
|
||||||
|
inp.skts.append(skt.unsqueeze(0))
|
||||||
|
|
||||||
|
arg = argparse.ArgumentParser()
|
||||||
|
|
||||||
|
arg.ref_mode = 'prevref'
|
||||||
|
arg.num_frames = num_frames
|
||||||
|
arg.num_steps = num_steps
|
||||||
|
arg.overlap = overlap
|
||||||
|
arg.prev_attn_steps = prev_attn_steps
|
||||||
|
arg.scale = [1.0, 1.0]
|
||||||
|
arg.seed = seed
|
||||||
|
arg.fps_id = fps_id
|
||||||
|
arg.motion_bucket_id = motion_bucket_id
|
||||||
|
arg.cond_aug = cond_aug
|
||||||
|
|
||||||
|
samples = sample_video(model, device, inp, arg, verbose=True)
|
||||||
|
|
||||||
|
return (LVCD_pipe, samples)
|
||||||
|
|
||||||
|
class LVCDDecoder:
|
||||||
|
@classmethod
|
||||||
|
def INPUT_TYPES(s):
|
||||||
|
return {
|
||||||
|
"required": {
|
||||||
|
"LVCD_pipe": ("LVCDPIPE",),
|
||||||
|
"samples": ("SVDSAMPLES",),
|
||||||
|
"decoding_t": ("INT", {"default": 10, "min": 1, "max": 100, "step": 1}),
|
||||||
|
"decoding_olap": ("INT", {"default": 3, "min": 0, "max": 100, "step": 1}),
|
||||||
|
"decoding_first": ("INT", {"default": 1, "min": 0, "max": 100, "step": 1}),
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
RETURN_TYPES = ("IMAGE",)
|
||||||
|
RETURN_NAMES = ("images", )
|
||||||
|
FUNCTION = "loadmodel"
|
||||||
|
CATEGORY = "ComfyUI-LVCDWrapper"
|
||||||
|
|
||||||
|
def loadmodel(self, LVCD_pipe, samples, decoding_t, decoding_olap, decoding_first):
|
||||||
|
|
||||||
|
device = mm.get_torch_device()
|
||||||
|
offload_device = mm.unet_offload_device()
|
||||||
|
mm.soft_empty_cache()
|
||||||
|
|
||||||
|
model = LVCD_pipe["model"]
|
||||||
|
|
||||||
|
arg = argparse.ArgumentParser()
|
||||||
|
|
||||||
|
arg.decoding_t = decoding_t
|
||||||
|
arg.decoding_olap = decoding_olap
|
||||||
|
arg.decoding_first = decoding_first
|
||||||
|
|
||||||
|
frames = decode_video(model, device, samples, arg)
|
||||||
|
|
||||||
|
min_value = frames.min()
|
||||||
|
max_value = frames.max()
|
||||||
|
|
||||||
|
frames = (frames - min_value) / (max_value - min_value)
|
||||||
|
frames = frames.permute(0, 2, 3, 1).cpu().float()
|
||||||
|
|
||||||
|
|
||||||
|
return (frames,)
|
||||||
|
|
||||||
|
NODE_CLASS_MAPPINGS = {
|
||||||
|
"LoadLVCDModel": LoadLVCDModel,
|
||||||
|
"LVCDSampler": LVCDSampler,
|
||||||
|
"LVCDDecoder": LVCDDecoder,
|
||||||
|
}
|
||||||
|
NODE_DISPLAY_NAME_MAPPINGS = {
|
||||||
|
"LoadLVCDModel": "Load LVCD Model",
|
||||||
|
"LVCDSampler": "LVCD Sampler",
|
||||||
|
"LVCDDecoder": "LVCD Decoder",
|
||||||
|
}
|
||||||
@@ -0,0 +1,4 @@
|
|||||||
|
omegaconf>=2.3.0
|
||||||
|
clip @ git+https://github.com/openai/CLIP.git
|
||||||
|
pytorch-lightning>=2.0.1
|
||||||
|
timm>=0.9.2
|
||||||
@@ -0,0 +1,4 @@
|
|||||||
|
from .models import AutoencodingEngine, DiffusionEngine
|
||||||
|
from .util import get_configs_path, instantiate_from_config
|
||||||
|
|
||||||
|
__version__ = "0.1.0"
|
||||||
@@ -0,0 +1,2 @@
|
|||||||
|
from .autoencoder import AutoencodingEngine
|
||||||
|
from .diffusion import DiffusionEngine
|
||||||
@@ -0,0 +1,615 @@
|
|||||||
|
import logging
|
||||||
|
import math
|
||||||
|
import re
|
||||||
|
from abc import abstractmethod
|
||||||
|
from contextlib import contextmanager
|
||||||
|
from typing import Any, Dict, List, Optional, Tuple, Union
|
||||||
|
|
||||||
|
import pytorch_lightning as pl
|
||||||
|
import torch
|
||||||
|
import torch.nn as nn
|
||||||
|
from einops import rearrange
|
||||||
|
from packaging import version
|
||||||
|
|
||||||
|
from ..modules.autoencoding.regularizers import AbstractRegularizer
|
||||||
|
from ..modules.ema import LitEma
|
||||||
|
from ..util import (default, get_nested_attribute, get_obj_from_str,
|
||||||
|
instantiate_from_config)
|
||||||
|
|
||||||
|
logpy = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
|
||||||
|
class AbstractAutoencoder(pl.LightningModule):
|
||||||
|
"""
|
||||||
|
This is the base class for all autoencoders, including image autoencoders, image autoencoders with discriminators,
|
||||||
|
unCLIP models, etc. Hence, it is fairly general, and specific features
|
||||||
|
(e.g. discriminator training, encoding, decoding) must be implemented in subclasses.
|
||||||
|
"""
|
||||||
|
|
||||||
|
def __init__(
|
||||||
|
self,
|
||||||
|
ema_decay: Union[None, float] = None,
|
||||||
|
monitor: Union[None, str] = None,
|
||||||
|
input_key: str = "jpg",
|
||||||
|
):
|
||||||
|
super().__init__()
|
||||||
|
|
||||||
|
self.input_key = input_key
|
||||||
|
self.use_ema = ema_decay is not None
|
||||||
|
if monitor is not None:
|
||||||
|
self.monitor = monitor
|
||||||
|
|
||||||
|
if self.use_ema:
|
||||||
|
self.model_ema = LitEma(self, decay=ema_decay)
|
||||||
|
logpy.info(f"Keeping EMAs of {len(list(self.model_ema.buffers()))}.")
|
||||||
|
|
||||||
|
if version.parse(torch.__version__) >= version.parse("2.0.0"):
|
||||||
|
self.automatic_optimization = False
|
||||||
|
|
||||||
|
def apply_ckpt(self, ckpt: Union[None, str, dict]):
|
||||||
|
if ckpt is None:
|
||||||
|
return
|
||||||
|
if isinstance(ckpt, str):
|
||||||
|
ckpt = {
|
||||||
|
"target": ".sgm.modules.checkpoint.CheckpointEngine",
|
||||||
|
"params": {"ckpt_path": ckpt},
|
||||||
|
}
|
||||||
|
engine = instantiate_from_config(ckpt)
|
||||||
|
engine(self)
|
||||||
|
|
||||||
|
@abstractmethod
|
||||||
|
def get_input(self, batch) -> Any:
|
||||||
|
raise NotImplementedError()
|
||||||
|
|
||||||
|
def on_train_batch_end(self, *args, **kwargs):
|
||||||
|
# for EMA computation
|
||||||
|
if self.use_ema:
|
||||||
|
self.model_ema(self)
|
||||||
|
|
||||||
|
@contextmanager
|
||||||
|
def ema_scope(self, context=None):
|
||||||
|
if self.use_ema:
|
||||||
|
self.model_ema.store(self.parameters())
|
||||||
|
self.model_ema.copy_to(self)
|
||||||
|
if context is not None:
|
||||||
|
logpy.info(f"{context}: Switched to EMA weights")
|
||||||
|
try:
|
||||||
|
yield None
|
||||||
|
finally:
|
||||||
|
if self.use_ema:
|
||||||
|
self.model_ema.restore(self.parameters())
|
||||||
|
if context is not None:
|
||||||
|
logpy.info(f"{context}: Restored training weights")
|
||||||
|
|
||||||
|
@abstractmethod
|
||||||
|
def encode(self, *args, **kwargs) -> torch.Tensor:
|
||||||
|
raise NotImplementedError("encode()-method of abstract base class called")
|
||||||
|
|
||||||
|
@abstractmethod
|
||||||
|
def decode(self, *args, **kwargs) -> torch.Tensor:
|
||||||
|
raise NotImplementedError("decode()-method of abstract base class called")
|
||||||
|
|
||||||
|
def instantiate_optimizer_from_config(self, params, lr, cfg):
|
||||||
|
logpy.info(f"loading >>> {cfg['target']} <<< optimizer from config")
|
||||||
|
return get_obj_from_str(cfg["target"])(
|
||||||
|
params, lr=lr, **cfg.get("params", dict())
|
||||||
|
)
|
||||||
|
|
||||||
|
def configure_optimizers(self) -> Any:
|
||||||
|
raise NotImplementedError()
|
||||||
|
|
||||||
|
|
||||||
|
class AutoencodingEngine(AbstractAutoencoder):
|
||||||
|
"""
|
||||||
|
Base class for all image autoencoders that we train, like VQGAN or AutoencoderKL
|
||||||
|
(we also restore them explicitly as special cases for legacy reasons).
|
||||||
|
Regularizations such as KL or VQ are moved to the regularizer class.
|
||||||
|
"""
|
||||||
|
|
||||||
|
def __init__(
|
||||||
|
self,
|
||||||
|
*args,
|
||||||
|
encoder_config: Dict,
|
||||||
|
decoder_config: Dict,
|
||||||
|
loss_config: Dict,
|
||||||
|
regularizer_config: Dict,
|
||||||
|
optimizer_config: Union[Dict, None] = None,
|
||||||
|
lr_g_factor: float = 1.0,
|
||||||
|
trainable_ae_params: Optional[List[List[str]]] = None,
|
||||||
|
ae_optimizer_args: Optional[List[dict]] = None,
|
||||||
|
trainable_disc_params: Optional[List[List[str]]] = None,
|
||||||
|
disc_optimizer_args: Optional[List[dict]] = None,
|
||||||
|
disc_start_iter: int = 0,
|
||||||
|
diff_boost_factor: float = 3.0,
|
||||||
|
ckpt_engine: Union[None, str, dict] = None,
|
||||||
|
ckpt_path: Optional[str] = None,
|
||||||
|
additional_decode_keys: Optional[List[str]] = None,
|
||||||
|
**kwargs,
|
||||||
|
):
|
||||||
|
super().__init__(*args, **kwargs)
|
||||||
|
self.automatic_optimization = False # pytorch lightning
|
||||||
|
|
||||||
|
self.encoder: torch.nn.Module = instantiate_from_config(encoder_config)
|
||||||
|
self.decoder: torch.nn.Module = instantiate_from_config(decoder_config)
|
||||||
|
self.loss: torch.nn.Module = instantiate_from_config(loss_config)
|
||||||
|
self.regularization: AbstractRegularizer = instantiate_from_config(
|
||||||
|
regularizer_config
|
||||||
|
)
|
||||||
|
self.optimizer_config = default(
|
||||||
|
optimizer_config, {"target": "torch.optim.Adam"}
|
||||||
|
)
|
||||||
|
self.diff_boost_factor = diff_boost_factor
|
||||||
|
self.disc_start_iter = disc_start_iter
|
||||||
|
self.lr_g_factor = lr_g_factor
|
||||||
|
self.trainable_ae_params = trainable_ae_params
|
||||||
|
if self.trainable_ae_params is not None:
|
||||||
|
self.ae_optimizer_args = default(
|
||||||
|
ae_optimizer_args,
|
||||||
|
[{} for _ in range(len(self.trainable_ae_params))],
|
||||||
|
)
|
||||||
|
assert len(self.ae_optimizer_args) == len(self.trainable_ae_params)
|
||||||
|
else:
|
||||||
|
self.ae_optimizer_args = [{}] # makes type consitent
|
||||||
|
|
||||||
|
self.trainable_disc_params = trainable_disc_params
|
||||||
|
if self.trainable_disc_params is not None:
|
||||||
|
self.disc_optimizer_args = default(
|
||||||
|
disc_optimizer_args,
|
||||||
|
[{} for _ in range(len(self.trainable_disc_params))],
|
||||||
|
)
|
||||||
|
assert len(self.disc_optimizer_args) == len(self.trainable_disc_params)
|
||||||
|
else:
|
||||||
|
self.disc_optimizer_args = [{}] # makes type consitent
|
||||||
|
|
||||||
|
if ckpt_path is not None:
|
||||||
|
assert ckpt_engine is None, "Can't set ckpt_engine and ckpt_path"
|
||||||
|
logpy.warn("Checkpoint path is deprecated, use `checkpoint_egnine` instead")
|
||||||
|
self.apply_ckpt(default(ckpt_path, ckpt_engine))
|
||||||
|
self.additional_decode_keys = set(default(additional_decode_keys, []))
|
||||||
|
|
||||||
|
def get_input(self, batch: Dict) -> torch.Tensor:
|
||||||
|
# assuming unified data format, dataloader returns a dict.
|
||||||
|
# image tensors should be scaled to -1 ... 1 and in channels-first
|
||||||
|
# format (e.g., bchw instead if bhwc)
|
||||||
|
return batch[self.input_key]
|
||||||
|
|
||||||
|
def get_autoencoder_params(self) -> list:
|
||||||
|
params = []
|
||||||
|
if hasattr(self.loss, "get_trainable_autoencoder_parameters"):
|
||||||
|
params += list(self.loss.get_trainable_autoencoder_parameters())
|
||||||
|
if hasattr(self.regularization, "get_trainable_parameters"):
|
||||||
|
params += list(self.regularization.get_trainable_parameters())
|
||||||
|
params = params + list(self.encoder.parameters())
|
||||||
|
params = params + list(self.decoder.parameters())
|
||||||
|
return params
|
||||||
|
|
||||||
|
def get_discriminator_params(self) -> list:
|
||||||
|
if hasattr(self.loss, "get_trainable_parameters"):
|
||||||
|
params = list(self.loss.get_trainable_parameters()) # e.g., discriminator
|
||||||
|
else:
|
||||||
|
params = []
|
||||||
|
return params
|
||||||
|
|
||||||
|
def get_last_layer(self):
|
||||||
|
return self.decoder.get_last_layer()
|
||||||
|
|
||||||
|
def encode(
|
||||||
|
self,
|
||||||
|
x: torch.Tensor,
|
||||||
|
return_reg_log: bool = False,
|
||||||
|
unregularized: bool = False,
|
||||||
|
) -> Union[torch.Tensor, Tuple[torch.Tensor, dict]]:
|
||||||
|
z = self.encoder(x)
|
||||||
|
if unregularized:
|
||||||
|
return z, dict()
|
||||||
|
z, reg_log = self.regularization(z)
|
||||||
|
if return_reg_log:
|
||||||
|
return z, reg_log
|
||||||
|
return z
|
||||||
|
|
||||||
|
def decode(self, z: torch.Tensor, **kwargs) -> torch.Tensor:
|
||||||
|
x = self.decoder(z, **kwargs)
|
||||||
|
return x
|
||||||
|
|
||||||
|
def forward(
|
||||||
|
self, x: torch.Tensor, **additional_decode_kwargs
|
||||||
|
) -> Tuple[torch.Tensor, torch.Tensor, dict]:
|
||||||
|
z, reg_log = self.encode(x, return_reg_log=True)
|
||||||
|
dec = self.decode(z, **additional_decode_kwargs)
|
||||||
|
return z, dec, reg_log
|
||||||
|
|
||||||
|
def inner_training_step(
|
||||||
|
self, batch: dict, batch_idx: int, optimizer_idx: int = 0
|
||||||
|
) -> torch.Tensor:
|
||||||
|
x = self.get_input(batch)
|
||||||
|
additional_decode_kwargs = {
|
||||||
|
key: batch[key] for key in self.additional_decode_keys.intersection(batch)
|
||||||
|
}
|
||||||
|
z, xrec, regularization_log = self(x, **additional_decode_kwargs)
|
||||||
|
if hasattr(self.loss, "forward_keys"):
|
||||||
|
extra_info = {
|
||||||
|
"z": z,
|
||||||
|
"optimizer_idx": optimizer_idx,
|
||||||
|
"global_step": self.global_step,
|
||||||
|
"last_layer": self.get_last_layer(),
|
||||||
|
"split": "train",
|
||||||
|
"regularization_log": regularization_log,
|
||||||
|
"autoencoder": self,
|
||||||
|
}
|
||||||
|
extra_info = {k: extra_info[k] for k in self.loss.forward_keys}
|
||||||
|
else:
|
||||||
|
extra_info = dict()
|
||||||
|
|
||||||
|
if optimizer_idx == 0:
|
||||||
|
# autoencode
|
||||||
|
out_loss = self.loss(x, xrec, **extra_info)
|
||||||
|
if isinstance(out_loss, tuple):
|
||||||
|
aeloss, log_dict_ae = out_loss
|
||||||
|
else:
|
||||||
|
# simple loss function
|
||||||
|
aeloss = out_loss
|
||||||
|
log_dict_ae = {"train/loss/rec": aeloss.detach()}
|
||||||
|
|
||||||
|
self.log_dict(
|
||||||
|
log_dict_ae,
|
||||||
|
prog_bar=False,
|
||||||
|
logger=True,
|
||||||
|
on_step=True,
|
||||||
|
on_epoch=True,
|
||||||
|
sync_dist=False,
|
||||||
|
)
|
||||||
|
self.log(
|
||||||
|
"loss",
|
||||||
|
aeloss.mean().detach(),
|
||||||
|
prog_bar=True,
|
||||||
|
logger=False,
|
||||||
|
on_epoch=False,
|
||||||
|
on_step=True,
|
||||||
|
)
|
||||||
|
return aeloss
|
||||||
|
elif optimizer_idx == 1:
|
||||||
|
# discriminator
|
||||||
|
discloss, log_dict_disc = self.loss(x, xrec, **extra_info)
|
||||||
|
# -> discriminator always needs to return a tuple
|
||||||
|
self.log_dict(
|
||||||
|
log_dict_disc, prog_bar=False, logger=True, on_step=True, on_epoch=True
|
||||||
|
)
|
||||||
|
return discloss
|
||||||
|
else:
|
||||||
|
raise NotImplementedError(f"Unknown optimizer {optimizer_idx}")
|
||||||
|
|
||||||
|
def training_step(self, batch: dict, batch_idx: int):
|
||||||
|
opts = self.optimizers()
|
||||||
|
if not isinstance(opts, list):
|
||||||
|
# Non-adversarial case
|
||||||
|
opts = [opts]
|
||||||
|
optimizer_idx = batch_idx % len(opts)
|
||||||
|
if self.global_step < self.disc_start_iter:
|
||||||
|
optimizer_idx = 0
|
||||||
|
opt = opts[optimizer_idx]
|
||||||
|
opt.zero_grad()
|
||||||
|
with opt.toggle_model():
|
||||||
|
loss = self.inner_training_step(
|
||||||
|
batch, batch_idx, optimizer_idx=optimizer_idx
|
||||||
|
)
|
||||||
|
self.manual_backward(loss)
|
||||||
|
opt.step()
|
||||||
|
|
||||||
|
def validation_step(self, batch: dict, batch_idx: int) -> Dict:
|
||||||
|
log_dict = self._validation_step(batch, batch_idx)
|
||||||
|
with self.ema_scope():
|
||||||
|
log_dict_ema = self._validation_step(batch, batch_idx, postfix="_ema")
|
||||||
|
log_dict.update(log_dict_ema)
|
||||||
|
return log_dict
|
||||||
|
|
||||||
|
def _validation_step(self, batch: dict, batch_idx: int, postfix: str = "") -> Dict:
|
||||||
|
x = self.get_input(batch)
|
||||||
|
|
||||||
|
z, xrec, regularization_log = self(x)
|
||||||
|
if hasattr(self.loss, "forward_keys"):
|
||||||
|
extra_info = {
|
||||||
|
"z": z,
|
||||||
|
"optimizer_idx": 0,
|
||||||
|
"global_step": self.global_step,
|
||||||
|
"last_layer": self.get_last_layer(),
|
||||||
|
"split": "val" + postfix,
|
||||||
|
"regularization_log": regularization_log,
|
||||||
|
"autoencoder": self,
|
||||||
|
}
|
||||||
|
extra_info = {k: extra_info[k] for k in self.loss.forward_keys}
|
||||||
|
else:
|
||||||
|
extra_info = dict()
|
||||||
|
out_loss = self.loss(x, xrec, **extra_info)
|
||||||
|
if isinstance(out_loss, tuple):
|
||||||
|
aeloss, log_dict_ae = out_loss
|
||||||
|
else:
|
||||||
|
# simple loss function
|
||||||
|
aeloss = out_loss
|
||||||
|
log_dict_ae = {f"val{postfix}/loss/rec": aeloss.detach()}
|
||||||
|
full_log_dict = log_dict_ae
|
||||||
|
|
||||||
|
if "optimizer_idx" in extra_info:
|
||||||
|
extra_info["optimizer_idx"] = 1
|
||||||
|
discloss, log_dict_disc = self.loss(x, xrec, **extra_info)
|
||||||
|
full_log_dict.update(log_dict_disc)
|
||||||
|
self.log(
|
||||||
|
f"val{postfix}/loss/rec",
|
||||||
|
log_dict_ae[f"val{postfix}/loss/rec"],
|
||||||
|
sync_dist=True,
|
||||||
|
)
|
||||||
|
self.log_dict(full_log_dict, sync_dist=True)
|
||||||
|
return full_log_dict
|
||||||
|
|
||||||
|
def get_param_groups(
|
||||||
|
self, parameter_names: List[List[str]], optimizer_args: List[dict]
|
||||||
|
) -> Tuple[List[Dict[str, Any]], int]:
|
||||||
|
groups = []
|
||||||
|
num_params = 0
|
||||||
|
for names, args in zip(parameter_names, optimizer_args):
|
||||||
|
params = []
|
||||||
|
for pattern_ in names:
|
||||||
|
pattern_params = []
|
||||||
|
pattern = re.compile(pattern_)
|
||||||
|
for p_name, param in self.named_parameters():
|
||||||
|
if re.match(pattern, p_name):
|
||||||
|
pattern_params.append(param)
|
||||||
|
num_params += param.numel()
|
||||||
|
if len(pattern_params) == 0:
|
||||||
|
logpy.warn(f"Did not find parameters for pattern {pattern_}")
|
||||||
|
params.extend(pattern_params)
|
||||||
|
groups.append({"params": params, **args})
|
||||||
|
return groups, num_params
|
||||||
|
|
||||||
|
def configure_optimizers(self) -> List[torch.optim.Optimizer]:
|
||||||
|
if self.trainable_ae_params is None:
|
||||||
|
ae_params = self.get_autoencoder_params()
|
||||||
|
else:
|
||||||
|
ae_params, num_ae_params = self.get_param_groups(
|
||||||
|
self.trainable_ae_params, self.ae_optimizer_args
|
||||||
|
)
|
||||||
|
logpy.info(f"Number of trainable autoencoder parameters: {num_ae_params:,}")
|
||||||
|
if self.trainable_disc_params is None:
|
||||||
|
disc_params = self.get_discriminator_params()
|
||||||
|
else:
|
||||||
|
disc_params, num_disc_params = self.get_param_groups(
|
||||||
|
self.trainable_disc_params, self.disc_optimizer_args
|
||||||
|
)
|
||||||
|
logpy.info(
|
||||||
|
f"Number of trainable discriminator parameters: {num_disc_params:,}"
|
||||||
|
)
|
||||||
|
opt_ae = self.instantiate_optimizer_from_config(
|
||||||
|
ae_params,
|
||||||
|
default(self.lr_g_factor, 1.0) * self.learning_rate,
|
||||||
|
self.optimizer_config,
|
||||||
|
)
|
||||||
|
opts = [opt_ae]
|
||||||
|
if len(disc_params) > 0:
|
||||||
|
opt_disc = self.instantiate_optimizer_from_config(
|
||||||
|
disc_params, self.learning_rate, self.optimizer_config
|
||||||
|
)
|
||||||
|
opts.append(opt_disc)
|
||||||
|
|
||||||
|
return opts
|
||||||
|
|
||||||
|
@torch.no_grad()
|
||||||
|
def log_images(
|
||||||
|
self, batch: dict, additional_log_kwargs: Optional[Dict] = None, **kwargs
|
||||||
|
) -> dict:
|
||||||
|
log = dict()
|
||||||
|
additional_decode_kwargs = {}
|
||||||
|
x = self.get_input(batch)
|
||||||
|
additional_decode_kwargs.update(
|
||||||
|
{key: batch[key] for key in self.additional_decode_keys.intersection(batch)}
|
||||||
|
)
|
||||||
|
|
||||||
|
_, xrec, _ = self(x, **additional_decode_kwargs)
|
||||||
|
log["inputs"] = x
|
||||||
|
log["reconstructions"] = xrec
|
||||||
|
diff = 0.5 * torch.abs(torch.clamp(xrec, -1.0, 1.0) - x)
|
||||||
|
diff.clamp_(0, 1.0)
|
||||||
|
log["diff"] = 2.0 * diff - 1.0
|
||||||
|
# diff_boost shows location of small errors, by boosting their
|
||||||
|
# brightness.
|
||||||
|
log["diff_boost"] = (
|
||||||
|
2.0 * torch.clamp(self.diff_boost_factor * diff, 0.0, 1.0) - 1
|
||||||
|
)
|
||||||
|
if hasattr(self.loss, "log_images"):
|
||||||
|
log.update(self.loss.log_images(x, xrec))
|
||||||
|
with self.ema_scope():
|
||||||
|
_, xrec_ema, _ = self(x, **additional_decode_kwargs)
|
||||||
|
log["reconstructions_ema"] = xrec_ema
|
||||||
|
diff_ema = 0.5 * torch.abs(torch.clamp(xrec_ema, -1.0, 1.0) - x)
|
||||||
|
diff_ema.clamp_(0, 1.0)
|
||||||
|
log["diff_ema"] = 2.0 * diff_ema - 1.0
|
||||||
|
log["diff_boost_ema"] = (
|
||||||
|
2.0 * torch.clamp(self.diff_boost_factor * diff_ema, 0.0, 1.0) - 1
|
||||||
|
)
|
||||||
|
if additional_log_kwargs:
|
||||||
|
additional_decode_kwargs.update(additional_log_kwargs)
|
||||||
|
_, xrec_add, _ = self(x, **additional_decode_kwargs)
|
||||||
|
log_str = "reconstructions-" + "-".join(
|
||||||
|
[f"{key}={additional_log_kwargs[key]}" for key in additional_log_kwargs]
|
||||||
|
)
|
||||||
|
log[log_str] = xrec_add
|
||||||
|
return log
|
||||||
|
|
||||||
|
|
||||||
|
class AutoencodingEngineLegacy(AutoencodingEngine):
|
||||||
|
def __init__(self, embed_dim: int, **kwargs):
|
||||||
|
self.max_batch_size = kwargs.pop("max_batch_size", None)
|
||||||
|
ddconfig = kwargs.pop("ddconfig")
|
||||||
|
ckpt_path = kwargs.pop("ckpt_path", None)
|
||||||
|
ckpt_engine = kwargs.pop("ckpt_engine", None)
|
||||||
|
super().__init__(
|
||||||
|
encoder_config={
|
||||||
|
"target": ".sgm.modules.diffusionmodules.model.Encoder",
|
||||||
|
"params": ddconfig,
|
||||||
|
},
|
||||||
|
decoder_config={
|
||||||
|
"target": ".sgm.modules.diffusionmodules.model.Decoder",
|
||||||
|
"params": ddconfig,
|
||||||
|
},
|
||||||
|
**kwargs,
|
||||||
|
)
|
||||||
|
self.quant_conv = torch.nn.Conv2d(
|
||||||
|
(1 + ddconfig["double_z"]) * ddconfig["z_channels"],
|
||||||
|
(1 + ddconfig["double_z"]) * embed_dim,
|
||||||
|
1,
|
||||||
|
)
|
||||||
|
self.post_quant_conv = torch.nn.Conv2d(embed_dim, ddconfig["z_channels"], 1)
|
||||||
|
self.embed_dim = embed_dim
|
||||||
|
|
||||||
|
self.apply_ckpt(default(ckpt_path, ckpt_engine))
|
||||||
|
|
||||||
|
def get_autoencoder_params(self) -> list:
|
||||||
|
params = super().get_autoencoder_params()
|
||||||
|
return params
|
||||||
|
|
||||||
|
def encode(
|
||||||
|
self, x: torch.Tensor, return_reg_log: bool = False
|
||||||
|
) -> Union[torch.Tensor, Tuple[torch.Tensor, dict]]:
|
||||||
|
if self.max_batch_size is None:
|
||||||
|
z = self.encoder(x)
|
||||||
|
z = self.quant_conv(z)
|
||||||
|
else:
|
||||||
|
N = x.shape[0]
|
||||||
|
bs = self.max_batch_size
|
||||||
|
n_batches = int(math.ceil(N / bs))
|
||||||
|
z = list()
|
||||||
|
for i_batch in range(n_batches):
|
||||||
|
z_batch = self.encoder(x[i_batch * bs : (i_batch + 1) * bs])
|
||||||
|
z_batch = self.quant_conv(z_batch)
|
||||||
|
z.append(z_batch)
|
||||||
|
z = torch.cat(z, 0)
|
||||||
|
|
||||||
|
z, reg_log = self.regularization(z)
|
||||||
|
if return_reg_log:
|
||||||
|
return z, reg_log
|
||||||
|
return z
|
||||||
|
|
||||||
|
def decode(self, z: torch.Tensor, **decoder_kwargs) -> torch.Tensor:
|
||||||
|
if self.max_batch_size is None:
|
||||||
|
dec = self.post_quant_conv(z)
|
||||||
|
dec = self.decoder(dec, **decoder_kwargs)
|
||||||
|
else:
|
||||||
|
N = z.shape[0]
|
||||||
|
bs = self.max_batch_size
|
||||||
|
n_batches = int(math.ceil(N / bs))
|
||||||
|
dec = list()
|
||||||
|
for i_batch in range(n_batches):
|
||||||
|
dec_batch = self.post_quant_conv(z[i_batch * bs : (i_batch + 1) * bs])
|
||||||
|
dec_batch = self.decoder(dec_batch, **decoder_kwargs)
|
||||||
|
dec.append(dec_batch)
|
||||||
|
dec = torch.cat(dec, 0)
|
||||||
|
|
||||||
|
return dec
|
||||||
|
|
||||||
|
|
||||||
|
class AutoencoderKL(AutoencodingEngineLegacy):
|
||||||
|
def __init__(self, **kwargs):
|
||||||
|
if "lossconfig" in kwargs:
|
||||||
|
kwargs["loss_config"] = kwargs.pop("lossconfig")
|
||||||
|
super().__init__(
|
||||||
|
regularizer_config={
|
||||||
|
"target": (
|
||||||
|
".sgm.modules.autoencoding.regularizers"
|
||||||
|
".DiagonalGaussianRegularizer"
|
||||||
|
)
|
||||||
|
},
|
||||||
|
**kwargs,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
class AutoencoderLegacyVQ(AutoencodingEngineLegacy):
|
||||||
|
def __init__(
|
||||||
|
self,
|
||||||
|
embed_dim: int,
|
||||||
|
n_embed: int,
|
||||||
|
sane_index_shape: bool = False,
|
||||||
|
**kwargs,
|
||||||
|
):
|
||||||
|
if "lossconfig" in kwargs:
|
||||||
|
logpy.warn(f"Parameter `lossconfig` is deprecated, use `loss_config`.")
|
||||||
|
kwargs["loss_config"] = kwargs.pop("lossconfig")
|
||||||
|
super().__init__(
|
||||||
|
regularizer_config={
|
||||||
|
"target": (
|
||||||
|
".sgm.modules.autoencoding.regularizers.quantize" ".VectorQuantizer"
|
||||||
|
),
|
||||||
|
"params": {
|
||||||
|
"n_e": n_embed,
|
||||||
|
"e_dim": embed_dim,
|
||||||
|
"sane_index_shape": sane_index_shape,
|
||||||
|
},
|
||||||
|
},
|
||||||
|
**kwargs,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
class IdentityFirstStage(AbstractAutoencoder):
|
||||||
|
def __init__(self, *args, **kwargs):
|
||||||
|
super().__init__(*args, **kwargs)
|
||||||
|
|
||||||
|
def get_input(self, x: Any) -> Any:
|
||||||
|
return x
|
||||||
|
|
||||||
|
def encode(self, x: Any, *args, **kwargs) -> Any:
|
||||||
|
return x
|
||||||
|
|
||||||
|
def decode(self, x: Any, *args, **kwargs) -> Any:
|
||||||
|
return x
|
||||||
|
|
||||||
|
|
||||||
|
class AEIntegerWrapper(nn.Module):
|
||||||
|
def __init__(
|
||||||
|
self,
|
||||||
|
model: nn.Module,
|
||||||
|
shape: Union[None, Tuple[int, int], List[int]] = (16, 16),
|
||||||
|
regularization_key: str = "regularization",
|
||||||
|
encoder_kwargs: Optional[Dict[str, Any]] = None,
|
||||||
|
):
|
||||||
|
super().__init__()
|
||||||
|
self.model = model
|
||||||
|
assert hasattr(model, "encode") and hasattr(
|
||||||
|
model, "decode"
|
||||||
|
), "Need AE interface"
|
||||||
|
self.regularization = get_nested_attribute(model, regularization_key)
|
||||||
|
self.shape = shape
|
||||||
|
self.encoder_kwargs = default(encoder_kwargs, {"return_reg_log": True})
|
||||||
|
|
||||||
|
def encode(self, x) -> torch.Tensor:
|
||||||
|
assert (
|
||||||
|
not self.training
|
||||||
|
), f"{self.__class__.__name__} only supports inference currently"
|
||||||
|
_, log = self.model.encode(x, **self.encoder_kwargs)
|
||||||
|
assert isinstance(log, dict)
|
||||||
|
inds = log["min_encoding_indices"]
|
||||||
|
return rearrange(inds, "b ... -> b (...)")
|
||||||
|
|
||||||
|
def decode(
|
||||||
|
self, inds: torch.Tensor, shape: Union[None, tuple, list] = None
|
||||||
|
) -> torch.Tensor:
|
||||||
|
# expect inds shape (b, s) with s = h*w
|
||||||
|
shape = default(shape, self.shape) # Optional[(h, w)]
|
||||||
|
if shape is not None:
|
||||||
|
assert len(shape) == 2, f"Unhandeled shape {shape}"
|
||||||
|
inds = rearrange(inds, "b (h w) -> b h w", h=shape[0], w=shape[1])
|
||||||
|
h = self.regularization.get_codebook_entry(inds) # (b, h, w, c)
|
||||||
|
h = rearrange(h, "b h w c -> b c h w")
|
||||||
|
return self.model.decode(h)
|
||||||
|
|
||||||
|
|
||||||
|
class AutoencoderKLModeOnly(AutoencodingEngineLegacy):
|
||||||
|
def __init__(self, **kwargs):
|
||||||
|
if "lossconfig" in kwargs:
|
||||||
|
kwargs["loss_config"] = kwargs.pop("lossconfig")
|
||||||
|
super().__init__(
|
||||||
|
regularizer_config={
|
||||||
|
"target": (
|
||||||
|
".sgm.modules.autoencoding.regularizers"
|
||||||
|
".DiagonalGaussianRegularizer"
|
||||||
|
),
|
||||||
|
"params": {"sample": False},
|
||||||
|
},
|
||||||
|
**kwargs,
|
||||||
|
)
|
||||||
@@ -0,0 +1,356 @@
|
|||||||
|
import math
|
||||||
|
from contextlib import contextmanager
|
||||||
|
from typing import Any, Dict, List, Optional, Tuple, Union
|
||||||
|
|
||||||
|
import pytorch_lightning as pl
|
||||||
|
import torch
|
||||||
|
from omegaconf import ListConfig, OmegaConf
|
||||||
|
from safetensors.torch import load_file as load_safetensors
|
||||||
|
from torch.optim.lr_scheduler import LambdaLR
|
||||||
|
from einops import rearrange
|
||||||
|
|
||||||
|
from ..modules import UNCONDITIONAL_CONFIG
|
||||||
|
from ..modules.autoencoding.temporal_ae import VideoDecoder
|
||||||
|
from ..modules.diffusionmodules.wrappers import OPENAIUNETWRAPPER
|
||||||
|
from ..modules.ema import LitEma
|
||||||
|
from ..util import (default, disabled_train, get_obj_from_str,
|
||||||
|
instantiate_from_config, log_txt_as_img)
|
||||||
|
|
||||||
|
|
||||||
|
class DiffusionEngine(pl.LightningModule):
|
||||||
|
def __init__(
|
||||||
|
self,
|
||||||
|
network_config,
|
||||||
|
denoiser_config,
|
||||||
|
first_stage_config,
|
||||||
|
conditioner_config: Union[None, Dict, ListConfig, OmegaConf] = None,
|
||||||
|
sampler_config: Union[None, Dict, ListConfig, OmegaConf] = None,
|
||||||
|
optimizer_config: Union[None, Dict, ListConfig, OmegaConf] = None,
|
||||||
|
scheduler_config: Union[None, Dict, ListConfig, OmegaConf] = None,
|
||||||
|
loss_fn_config: Union[None, Dict, ListConfig, OmegaConf] = None,
|
||||||
|
network_wrapper: Union[None, str] = None,
|
||||||
|
ckpt_path: Union[None, str] = None,
|
||||||
|
use_ema: bool = False,
|
||||||
|
ema_decay_rate: float = 0.9999,
|
||||||
|
scale_factor: float = 1.0,
|
||||||
|
disable_first_stage_autocast=False,
|
||||||
|
input_key: str = "jpg",
|
||||||
|
log_keys: Union[List, None] = None,
|
||||||
|
no_cond_log: bool = False,
|
||||||
|
compile_model: bool = False,
|
||||||
|
en_and_decode_n_samples_a_time: Optional[int] = None,
|
||||||
|
):
|
||||||
|
super().__init__()
|
||||||
|
self.log_keys = log_keys
|
||||||
|
self.input_key = input_key
|
||||||
|
self.optimizer_config = default(
|
||||||
|
optimizer_config, {"target": "torch.optim.AdamW"}
|
||||||
|
)
|
||||||
|
model = instantiate_from_config(network_config)
|
||||||
|
self.model = get_obj_from_str(default(network_wrapper, OPENAIUNETWRAPPER))(
|
||||||
|
model, compile_model=compile_model
|
||||||
|
)
|
||||||
|
|
||||||
|
self.denoiser = instantiate_from_config(denoiser_config)
|
||||||
|
self.sampler = (
|
||||||
|
instantiate_from_config(sampler_config)
|
||||||
|
if sampler_config is not None
|
||||||
|
else None
|
||||||
|
)
|
||||||
|
self.conditioner = instantiate_from_config(
|
||||||
|
default(conditioner_config, UNCONDITIONAL_CONFIG)
|
||||||
|
)
|
||||||
|
self.scheduler_config = scheduler_config
|
||||||
|
self._init_first_stage(first_stage_config)
|
||||||
|
|
||||||
|
self.loss_fn = (
|
||||||
|
instantiate_from_config(loss_fn_config)
|
||||||
|
if loss_fn_config is not None
|
||||||
|
else None
|
||||||
|
)
|
||||||
|
|
||||||
|
self.use_ema = use_ema
|
||||||
|
if self.use_ema:
|
||||||
|
self.model_ema = LitEma(self.model, decay=ema_decay_rate)
|
||||||
|
print(f"Keeping EMAs of {len(list(self.model_ema.buffers()))}.")
|
||||||
|
|
||||||
|
self.scale_factor = scale_factor
|
||||||
|
self.disable_first_stage_autocast = disable_first_stage_autocast
|
||||||
|
self.no_cond_log = no_cond_log
|
||||||
|
|
||||||
|
if ckpt_path is not None:
|
||||||
|
self.init_from_ckpt(ckpt_path)
|
||||||
|
|
||||||
|
self.en_and_decode_n_samples_a_time = en_and_decode_n_samples_a_time
|
||||||
|
|
||||||
|
def init_from_ckpt(
|
||||||
|
self,
|
||||||
|
path: str,
|
||||||
|
) -> None:
|
||||||
|
if path.endswith("ckpt"):
|
||||||
|
sd = torch.load(path, map_location="cpu")["state_dict"]
|
||||||
|
elif path.endswith("safetensors"):
|
||||||
|
sd = load_safetensors(path)
|
||||||
|
else:
|
||||||
|
raise NotImplementedError
|
||||||
|
|
||||||
|
missing, unexpected = self.load_state_dict(sd, strict=False)
|
||||||
|
# print(
|
||||||
|
# f"Restored from {path} with {len(missing)} missing and {len(unexpected)} unexpected keys"
|
||||||
|
# )
|
||||||
|
# if len(missing) > 0:
|
||||||
|
# print(f"Missing Keys: {missing}")
|
||||||
|
# if len(unexpected) > 0:
|
||||||
|
# print(f"Unexpected Keys: {unexpected}")
|
||||||
|
|
||||||
|
def _init_first_stage(self, config):
|
||||||
|
model = instantiate_from_config(config)#.eval()
|
||||||
|
# Train function is overwritten by a empty function to ensure no training
|
||||||
|
model.train = disabled_train
|
||||||
|
for param in model.parameters():
|
||||||
|
param.requires_grad = False
|
||||||
|
self.first_stage_model = model
|
||||||
|
|
||||||
|
def get_input(self, batch):
|
||||||
|
# assuming unified data format, dataloader returns a dict.
|
||||||
|
# image tensors should be scaled to -1 ... 1 and in bchw format
|
||||||
|
if 'num_video_frames' in batch.keys():
|
||||||
|
for k in batch.keys():
|
||||||
|
if k not in ['num_video_frames', 'image_only_indicator', 'first_sigma']:
|
||||||
|
batch[k] = rearrange(batch[k], 'b t ... -> (b t) ...')
|
||||||
|
elif k in ['num_video_frames', 'first_sigma']:
|
||||||
|
batch[k] = batch[k].detach().cpu().numpy()[0]
|
||||||
|
return batch[self.input_key]
|
||||||
|
|
||||||
|
@torch.no_grad()
|
||||||
|
def decode_first_stage(self, z):
|
||||||
|
z = 1.0 / self.scale_factor * z
|
||||||
|
n_samples = default(self.en_and_decode_n_samples_a_time, z.shape[0])
|
||||||
|
|
||||||
|
n_rounds = math.ceil(z.shape[0] / n_samples)
|
||||||
|
all_out = []
|
||||||
|
with torch.autocast("cuda", enabled=not self.disable_first_stage_autocast):
|
||||||
|
for n in range(n_rounds):
|
||||||
|
if isinstance(self.first_stage_model.decoder, VideoDecoder):
|
||||||
|
kwargs = {"timesteps": len(z[n * n_samples : (n + 1) * n_samples])}
|
||||||
|
else:
|
||||||
|
kwargs = {}
|
||||||
|
out = self.first_stage_model.decode(
|
||||||
|
z[n * n_samples : (n + 1) * n_samples], **kwargs
|
||||||
|
)
|
||||||
|
all_out.append(out)
|
||||||
|
out = torch.cat(all_out, dim=0)
|
||||||
|
return out
|
||||||
|
|
||||||
|
@torch.no_grad()
|
||||||
|
def encode_first_stage(self, x):
|
||||||
|
n_samples = default(self.en_and_decode_n_samples_a_time, x.shape[0])
|
||||||
|
n_samples = 1
|
||||||
|
n_rounds = math.ceil(x.shape[0] / n_samples)
|
||||||
|
all_out = []
|
||||||
|
with torch.autocast("cuda", enabled=not self.disable_first_stage_autocast):
|
||||||
|
for n in range(n_rounds):
|
||||||
|
out = self.first_stage_model.encode( x[n * n_samples : (n + 1) * n_samples] )
|
||||||
|
self.first_stage_model.zero_grad(set_to_none=True)
|
||||||
|
torch.cuda.empty_cache()
|
||||||
|
all_out.append(out)
|
||||||
|
|
||||||
|
z = torch.cat(all_out, dim=0)
|
||||||
|
z = self.scale_factor * z
|
||||||
|
return z
|
||||||
|
|
||||||
|
def forward(self, x, batch):
|
||||||
|
loss = self.loss_fn(self.model, self.denoiser, self.conditioner, x, batch)
|
||||||
|
loss_mean = loss.mean()
|
||||||
|
loss_dict = {"loss": loss_mean}
|
||||||
|
return loss_mean, loss_dict
|
||||||
|
|
||||||
|
def shared_step(self, batch: Dict) -> Any:
|
||||||
|
x = self.get_input(batch)
|
||||||
|
x = self.encode_first_stage(x)
|
||||||
|
batch["global_step"] = self.global_step
|
||||||
|
loss, loss_dict = self(x, batch)
|
||||||
|
return loss, loss_dict
|
||||||
|
|
||||||
|
def training_step(self, batch, batch_idx):
|
||||||
|
loss, loss_dict = self.shared_step(batch)
|
||||||
|
|
||||||
|
self.log_dict(
|
||||||
|
loss_dict, prog_bar=True, logger=True, on_step=True, on_epoch=False
|
||||||
|
)
|
||||||
|
|
||||||
|
self.log(
|
||||||
|
"global_step",
|
||||||
|
self.global_step,
|
||||||
|
prog_bar=True,
|
||||||
|
logger=True,
|
||||||
|
on_step=True,
|
||||||
|
on_epoch=False,
|
||||||
|
)
|
||||||
|
|
||||||
|
if self.scheduler_config is not None:
|
||||||
|
lr = self.optimizers().param_groups[0]["lr"]
|
||||||
|
self.log(
|
||||||
|
"lr_abs", lr, prog_bar=True, logger=True, on_step=True, on_epoch=False
|
||||||
|
)
|
||||||
|
|
||||||
|
return loss
|
||||||
|
|
||||||
|
def on_train_start(self, *args, **kwargs):
|
||||||
|
if self.sampler is None or self.loss_fn is None:
|
||||||
|
raise ValueError("Sampler and loss function need to be set for training.")
|
||||||
|
|
||||||
|
def on_train_batch_end(self, *args, **kwargs):
|
||||||
|
if self.use_ema:
|
||||||
|
self.model_ema(self.model)
|
||||||
|
|
||||||
|
@contextmanager
|
||||||
|
def ema_scope(self, context=None):
|
||||||
|
if self.use_ema:
|
||||||
|
self.model_ema.store(self.model.parameters())
|
||||||
|
self.model_ema.copy_to(self.model)
|
||||||
|
if context is not None:
|
||||||
|
print(f"{context}: Switched to EMA weights")
|
||||||
|
try:
|
||||||
|
yield None
|
||||||
|
finally:
|
||||||
|
if self.use_ema:
|
||||||
|
self.model_ema.restore(self.model.parameters())
|
||||||
|
if context is not None:
|
||||||
|
print(f"{context}: Restored training weights")
|
||||||
|
|
||||||
|
def instantiate_optimizer_from_config(self, params, lr, cfg):
|
||||||
|
return get_obj_from_str(cfg["target"])(
|
||||||
|
params, lr=lr, **cfg.get("params", dict())
|
||||||
|
)
|
||||||
|
|
||||||
|
def configure_optimizers(self):
|
||||||
|
lr = self.learning_rate
|
||||||
|
params = list(self.model.parameters())
|
||||||
|
for embedder in self.conditioner.embedders:
|
||||||
|
if embedder.is_trainable:
|
||||||
|
params = params + list(embedder.parameters())
|
||||||
|
opt = self.instantiate_optimizer_from_config(params, lr, self.optimizer_config)
|
||||||
|
if self.scheduler_config is not None:
|
||||||
|
scheduler = instantiate_from_config(self.scheduler_config)
|
||||||
|
print("Setting up LambdaLR scheduler...")
|
||||||
|
scheduler = [
|
||||||
|
{
|
||||||
|
"scheduler": LambdaLR(opt, lr_lambda=scheduler.schedule),
|
||||||
|
"interval": "step",
|
||||||
|
"frequency": 1,
|
||||||
|
}
|
||||||
|
]
|
||||||
|
return [opt], scheduler
|
||||||
|
return opt
|
||||||
|
|
||||||
|
@torch.no_grad()
|
||||||
|
def sample(
|
||||||
|
self,
|
||||||
|
cond: Dict,
|
||||||
|
uc: Union[Dict, None] = None,
|
||||||
|
batch_size: int = 16,
|
||||||
|
shape: Union[None, Tuple, List] = None,
|
||||||
|
**kwargs,
|
||||||
|
):
|
||||||
|
randn = torch.randn(batch_size, *shape).to(self.device)
|
||||||
|
|
||||||
|
denoiser = lambda input, sigma, c: self.denoiser(
|
||||||
|
self.model, input, sigma, c, **kwargs
|
||||||
|
)
|
||||||
|
samples = self.sampler(denoiser, randn, cond, uc=uc)
|
||||||
|
return samples
|
||||||
|
|
||||||
|
@torch.no_grad()
|
||||||
|
def log_conditionings(self, batch: Dict, n: int) -> Dict:
|
||||||
|
"""
|
||||||
|
Defines heuristics to log different conditionings.
|
||||||
|
These can be lists of strings (text-to-image), tensors, ints, ...
|
||||||
|
"""
|
||||||
|
image_h, image_w = batch[self.input_key].shape[2:]
|
||||||
|
log = dict()
|
||||||
|
|
||||||
|
for embedder in self.conditioner.embedders:
|
||||||
|
if (
|
||||||
|
(self.log_keys is None) or (embedder.input_key in self.log_keys)
|
||||||
|
) and not self.no_cond_log:
|
||||||
|
x = batch[embedder.input_key][:n]
|
||||||
|
if isinstance(x, torch.Tensor):
|
||||||
|
if x.dim() == 1:
|
||||||
|
# class-conditional, convert integer to string
|
||||||
|
x = [str(x[i].detach().cpu().numpy()) for i in range(x.shape[0])]
|
||||||
|
xc = log_txt_as_img((image_h, image_w), x, size=image_h // 4)
|
||||||
|
elif x.dim() == 2:
|
||||||
|
# size and crop cond and the like
|
||||||
|
x = [
|
||||||
|
"x".join([str(xx) for xx in x[i].tolist()])
|
||||||
|
for i in range(x.shape[0])
|
||||||
|
]
|
||||||
|
xc = log_txt_as_img((image_h, image_w), x, size=image_h // 20)
|
||||||
|
else:
|
||||||
|
raise NotImplementedError()
|
||||||
|
elif isinstance(x, (List, ListConfig)):
|
||||||
|
if isinstance(x[0], str):
|
||||||
|
# strings
|
||||||
|
xc = log_txt_as_img((image_h, image_w), x, size=image_h // 20)
|
||||||
|
else:
|
||||||
|
raise NotImplementedError()
|
||||||
|
else:
|
||||||
|
raise NotImplementedError()
|
||||||
|
log[embedder.input_key] = xc
|
||||||
|
return log
|
||||||
|
|
||||||
|
@torch.no_grad()
|
||||||
|
def log_images(
|
||||||
|
self,
|
||||||
|
batch: Dict,
|
||||||
|
N: int = 8,
|
||||||
|
sample: bool = True,
|
||||||
|
ucg_keys: List[str] = None,
|
||||||
|
**kwargs,
|
||||||
|
) -> Dict:
|
||||||
|
conditioner_input_keys = [e.input_key for e in self.conditioner.embedders]
|
||||||
|
if ucg_keys:
|
||||||
|
assert all(map(lambda x: x in conditioner_input_keys, ucg_keys)), (
|
||||||
|
"Each defined ucg key for sampling must be in the provided conditioner input keys,"
|
||||||
|
f"but we have {ucg_keys} vs. {conditioner_input_keys}"
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
ucg_keys = conditioner_input_keys
|
||||||
|
log = dict()
|
||||||
|
|
||||||
|
torch.cuda.empty_cache()
|
||||||
|
|
||||||
|
x = self.get_input(batch)
|
||||||
|
|
||||||
|
c, uc = self.conditioner.get_unconditional_conditioning(
|
||||||
|
batch,
|
||||||
|
force_uc_zero_embeddings=ucg_keys
|
||||||
|
if len(self.conditioner.embedders) > 0
|
||||||
|
else [],
|
||||||
|
)
|
||||||
|
|
||||||
|
sampling_kwargs = {}
|
||||||
|
|
||||||
|
N = min(x.shape[0], N)
|
||||||
|
x = x.to(self.device)[:N]
|
||||||
|
log["inputs"] = x
|
||||||
|
|
||||||
|
torch.cuda.empty_cache()
|
||||||
|
|
||||||
|
z = self.encode_first_stage(x)
|
||||||
|
log["reconstructions"] = self.decode_first_stage(z)
|
||||||
|
log.update(self.log_conditionings(batch, N))
|
||||||
|
|
||||||
|
for k in c:
|
||||||
|
if isinstance(c[k], torch.Tensor):
|
||||||
|
c[k], uc[k] = map(lambda y: y[k][:N].to(self.device), (c, uc))
|
||||||
|
|
||||||
|
if sample:
|
||||||
|
with self.ema_scope("Plotting"):
|
||||||
|
samples = self.sample(
|
||||||
|
c, shape=z.shape[1:], uc=uc, batch_size=N, **sampling_kwargs
|
||||||
|
)
|
||||||
|
samples = self.decode_first_stage(samples)
|
||||||
|
log["samples"] = samples
|
||||||
|
return log
|
||||||
@@ -0,0 +1,6 @@
|
|||||||
|
from .encoders.modules import GeneralConditioner
|
||||||
|
|
||||||
|
UNCONDITIONAL_CONFIG = {
|
||||||
|
"target": ".sgm.modules.GeneralConditioner",
|
||||||
|
"params": {"emb_models": []},
|
||||||
|
}
|
||||||
@@ -0,0 +1,7 @@
|
|||||||
|
__all__ = [
|
||||||
|
"GeneralLPIPSWithDiscriminator",
|
||||||
|
"LatentLPIPS",
|
||||||
|
]
|
||||||
|
|
||||||
|
from .discriminator_loss import GeneralLPIPSWithDiscriminator
|
||||||
|
from .lpips import LatentLPIPS
|
||||||
@@ -0,0 +1,306 @@
|
|||||||
|
from typing import Dict, Iterator, List, Optional, Tuple, Union
|
||||||
|
|
||||||
|
import numpy as np
|
||||||
|
import torch
|
||||||
|
import torch.nn as nn
|
||||||
|
import torchvision
|
||||||
|
from einops import rearrange
|
||||||
|
from matplotlib import colormaps
|
||||||
|
from matplotlib import pyplot as plt
|
||||||
|
|
||||||
|
from ....util import default, instantiate_from_config
|
||||||
|
from ..lpips.loss.lpips import LPIPS
|
||||||
|
from ..lpips.model.model import weights_init
|
||||||
|
from ..lpips.vqperceptual import hinge_d_loss, vanilla_d_loss
|
||||||
|
|
||||||
|
|
||||||
|
class GeneralLPIPSWithDiscriminator(nn.Module):
|
||||||
|
def __init__(
|
||||||
|
self,
|
||||||
|
disc_start: int,
|
||||||
|
logvar_init: float = 0.0,
|
||||||
|
disc_num_layers: int = 3,
|
||||||
|
disc_in_channels: int = 3,
|
||||||
|
disc_factor: float = 1.0,
|
||||||
|
disc_weight: float = 1.0,
|
||||||
|
perceptual_weight: float = 1.0,
|
||||||
|
disc_loss: str = "hinge",
|
||||||
|
scale_input_to_tgt_size: bool = False,
|
||||||
|
dims: int = 2,
|
||||||
|
learn_logvar: bool = False,
|
||||||
|
regularization_weights: Union[None, Dict[str, float]] = None,
|
||||||
|
additional_log_keys: Optional[List[str]] = None,
|
||||||
|
discriminator_config: Optional[Dict] = None,
|
||||||
|
):
|
||||||
|
super().__init__()
|
||||||
|
self.dims = dims
|
||||||
|
if self.dims > 2:
|
||||||
|
print(
|
||||||
|
f"running with dims={dims}. This means that for perceptual loss "
|
||||||
|
f"calculation, the LPIPS loss will be applied to each frame "
|
||||||
|
f"independently."
|
||||||
|
)
|
||||||
|
self.scale_input_to_tgt_size = scale_input_to_tgt_size
|
||||||
|
assert disc_loss in ["hinge", "vanilla"]
|
||||||
|
self.perceptual_loss = LPIPS().eval()
|
||||||
|
self.perceptual_weight = perceptual_weight
|
||||||
|
# output log variance
|
||||||
|
self.logvar = nn.Parameter(
|
||||||
|
torch.full((), logvar_init), requires_grad=learn_logvar
|
||||||
|
)
|
||||||
|
self.learn_logvar = learn_logvar
|
||||||
|
|
||||||
|
discriminator_config = default(
|
||||||
|
discriminator_config,
|
||||||
|
{
|
||||||
|
"target": ".sgm.modules.autoencoding.lpips.model.model.NLayerDiscriminator",
|
||||||
|
"params": {
|
||||||
|
"input_nc": disc_in_channels,
|
||||||
|
"n_layers": disc_num_layers,
|
||||||
|
"use_actnorm": False,
|
||||||
|
},
|
||||||
|
},
|
||||||
|
)
|
||||||
|
|
||||||
|
self.discriminator = instantiate_from_config(discriminator_config).apply(
|
||||||
|
weights_init
|
||||||
|
)
|
||||||
|
self.discriminator_iter_start = disc_start
|
||||||
|
self.disc_loss = hinge_d_loss if disc_loss == "hinge" else vanilla_d_loss
|
||||||
|
self.disc_factor = disc_factor
|
||||||
|
self.discriminator_weight = disc_weight
|
||||||
|
self.regularization_weights = default(regularization_weights, {})
|
||||||
|
|
||||||
|
self.forward_keys = [
|
||||||
|
"optimizer_idx",
|
||||||
|
"global_step",
|
||||||
|
"last_layer",
|
||||||
|
"split",
|
||||||
|
"regularization_log",
|
||||||
|
]
|
||||||
|
|
||||||
|
self.additional_log_keys = set(default(additional_log_keys, []))
|
||||||
|
self.additional_log_keys.update(set(self.regularization_weights.keys()))
|
||||||
|
|
||||||
|
def get_trainable_parameters(self) -> Iterator[nn.Parameter]:
|
||||||
|
return self.discriminator.parameters()
|
||||||
|
|
||||||
|
def get_trainable_autoencoder_parameters(self) -> Iterator[nn.Parameter]:
|
||||||
|
if self.learn_logvar:
|
||||||
|
yield self.logvar
|
||||||
|
yield from ()
|
||||||
|
|
||||||
|
@torch.no_grad()
|
||||||
|
def log_images(
|
||||||
|
self, inputs: torch.Tensor, reconstructions: torch.Tensor
|
||||||
|
) -> Dict[str, torch.Tensor]:
|
||||||
|
# calc logits of real/fake
|
||||||
|
logits_real = self.discriminator(inputs.contiguous().detach())
|
||||||
|
if len(logits_real.shape) < 4:
|
||||||
|
# Non patch-discriminator
|
||||||
|
return dict()
|
||||||
|
logits_fake = self.discriminator(reconstructions.contiguous().detach())
|
||||||
|
# -> (b, 1, h, w)
|
||||||
|
|
||||||
|
# parameters for colormapping
|
||||||
|
high = max(logits_fake.abs().max(), logits_real.abs().max()).item()
|
||||||
|
cmap = colormaps["PiYG"] # diverging colormap
|
||||||
|
|
||||||
|
def to_colormap(logits: torch.Tensor) -> torch.Tensor:
|
||||||
|
"""(b, 1, ...) -> (b, 3, ...)"""
|
||||||
|
logits = (logits + high) / (2 * high)
|
||||||
|
logits_np = cmap(logits.cpu().numpy())[..., :3] # truncate alpha channel
|
||||||
|
# -> (b, 1, ..., 3)
|
||||||
|
logits = torch.from_numpy(logits_np).to(logits.device)
|
||||||
|
return rearrange(logits, "b 1 ... c -> b c ...")
|
||||||
|
|
||||||
|
logits_real = torch.nn.functional.interpolate(
|
||||||
|
logits_real,
|
||||||
|
size=inputs.shape[-2:],
|
||||||
|
mode="nearest",
|
||||||
|
antialias=False,
|
||||||
|
)
|
||||||
|
logits_fake = torch.nn.functional.interpolate(
|
||||||
|
logits_fake,
|
||||||
|
size=reconstructions.shape[-2:],
|
||||||
|
mode="nearest",
|
||||||
|
antialias=False,
|
||||||
|
)
|
||||||
|
|
||||||
|
# alpha value of logits for overlay
|
||||||
|
alpha_real = torch.abs(logits_real) / high
|
||||||
|
alpha_fake = torch.abs(logits_fake) / high
|
||||||
|
# -> (b, 1, h, w) in range [0, 0.5]
|
||||||
|
# alpha value of lines don't really matter, since the values are the same
|
||||||
|
# for both images and logits anyway
|
||||||
|
grid_alpha_real = torchvision.utils.make_grid(alpha_real, nrow=4)
|
||||||
|
grid_alpha_fake = torchvision.utils.make_grid(alpha_fake, nrow=4)
|
||||||
|
grid_alpha = 0.8 * torch.cat((grid_alpha_real, grid_alpha_fake), dim=1)
|
||||||
|
# -> (1, h, w)
|
||||||
|
# blend logits and images together
|
||||||
|
|
||||||
|
# prepare logits for plotting
|
||||||
|
logits_real = to_colormap(logits_real)
|
||||||
|
logits_fake = to_colormap(logits_fake)
|
||||||
|
# resize logits
|
||||||
|
# -> (b, 3, h, w)
|
||||||
|
|
||||||
|
# make some grids
|
||||||
|
# add all logits to one plot
|
||||||
|
logits_real = torchvision.utils.make_grid(logits_real, nrow=4)
|
||||||
|
logits_fake = torchvision.utils.make_grid(logits_fake, nrow=4)
|
||||||
|
# I just love how torchvision calls the number of columns `nrow`
|
||||||
|
grid_logits = torch.cat((logits_real, logits_fake), dim=1)
|
||||||
|
# -> (3, h, w)
|
||||||
|
|
||||||
|
grid_images_real = torchvision.utils.make_grid(0.5 * inputs + 0.5, nrow=4)
|
||||||
|
grid_images_fake = torchvision.utils.make_grid(
|
||||||
|
0.5 * reconstructions + 0.5, nrow=4
|
||||||
|
)
|
||||||
|
grid_images = torch.cat((grid_images_real, grid_images_fake), dim=1)
|
||||||
|
# -> (3, h, w) in range [0, 1]
|
||||||
|
|
||||||
|
grid_blend = grid_alpha * grid_logits + (1 - grid_alpha) * grid_images
|
||||||
|
|
||||||
|
# Create labeled colorbar
|
||||||
|
dpi = 100
|
||||||
|
height = 128 / dpi
|
||||||
|
width = grid_logits.shape[2] / dpi
|
||||||
|
fig, ax = plt.subplots(figsize=(width, height), dpi=dpi)
|
||||||
|
img = ax.imshow(np.array([[-high, high]]), cmap=cmap)
|
||||||
|
plt.colorbar(
|
||||||
|
img,
|
||||||
|
cax=ax,
|
||||||
|
orientation="horizontal",
|
||||||
|
fraction=0.9,
|
||||||
|
aspect=width / height,
|
||||||
|
pad=0.0,
|
||||||
|
)
|
||||||
|
img.set_visible(False)
|
||||||
|
fig.tight_layout()
|
||||||
|
fig.canvas.draw()
|
||||||
|
# manually convert figure to numpy
|
||||||
|
cbar_np = np.frombuffer(fig.canvas.tostring_rgb(), dtype=np.uint8)
|
||||||
|
cbar_np = cbar_np.reshape(fig.canvas.get_width_height()[::-1] + (3,))
|
||||||
|
cbar = torch.from_numpy(cbar_np.copy()).to(grid_logits.dtype) / 255.0
|
||||||
|
cbar = rearrange(cbar, "h w c -> c h w").to(grid_logits.device)
|
||||||
|
|
||||||
|
# Add colorbar to plot
|
||||||
|
annotated_grid = torch.cat((grid_logits, cbar), dim=1)
|
||||||
|
blended_grid = torch.cat((grid_blend, cbar), dim=1)
|
||||||
|
return {
|
||||||
|
"vis_logits": 2 * annotated_grid[None, ...] - 1,
|
||||||
|
"vis_logits_blended": 2 * blended_grid[None, ...] - 1,
|
||||||
|
}
|
||||||
|
|
||||||
|
def calculate_adaptive_weight(
|
||||||
|
self, nll_loss: torch.Tensor, g_loss: torch.Tensor, last_layer: torch.Tensor
|
||||||
|
) -> torch.Tensor:
|
||||||
|
nll_grads = torch.autograd.grad(nll_loss, last_layer, retain_graph=True)[0]
|
||||||
|
g_grads = torch.autograd.grad(g_loss, last_layer, retain_graph=True)[0]
|
||||||
|
|
||||||
|
d_weight = torch.norm(nll_grads) / (torch.norm(g_grads) + 1e-4)
|
||||||
|
d_weight = torch.clamp(d_weight, 0.0, 1e4).detach()
|
||||||
|
d_weight = d_weight * self.discriminator_weight
|
||||||
|
return d_weight
|
||||||
|
|
||||||
|
def forward(
|
||||||
|
self,
|
||||||
|
inputs: torch.Tensor,
|
||||||
|
reconstructions: torch.Tensor,
|
||||||
|
*, # added because I changed the order here
|
||||||
|
regularization_log: Dict[str, torch.Tensor],
|
||||||
|
optimizer_idx: int,
|
||||||
|
global_step: int,
|
||||||
|
last_layer: torch.Tensor,
|
||||||
|
split: str = "train",
|
||||||
|
weights: Union[None, float, torch.Tensor] = None,
|
||||||
|
) -> Tuple[torch.Tensor, dict]:
|
||||||
|
if self.scale_input_to_tgt_size:
|
||||||
|
inputs = torch.nn.functional.interpolate(
|
||||||
|
inputs, reconstructions.shape[2:], mode="bicubic", antialias=True
|
||||||
|
)
|
||||||
|
|
||||||
|
if self.dims > 2:
|
||||||
|
inputs, reconstructions = map(
|
||||||
|
lambda x: rearrange(x, "b c t h w -> (b t) c h w"),
|
||||||
|
(inputs, reconstructions),
|
||||||
|
)
|
||||||
|
|
||||||
|
rec_loss = torch.abs(inputs.contiguous() - reconstructions.contiguous())
|
||||||
|
if self.perceptual_weight > 0:
|
||||||
|
p_loss = self.perceptual_loss(
|
||||||
|
inputs.contiguous(), reconstructions.contiguous()
|
||||||
|
)
|
||||||
|
rec_loss = rec_loss + self.perceptual_weight * p_loss
|
||||||
|
|
||||||
|
nll_loss, weighted_nll_loss = self.get_nll_loss(rec_loss, weights)
|
||||||
|
|
||||||
|
# now the GAN part
|
||||||
|
if optimizer_idx == 0:
|
||||||
|
# generator update
|
||||||
|
if global_step >= self.discriminator_iter_start or not self.training:
|
||||||
|
logits_fake = self.discriminator(reconstructions.contiguous())
|
||||||
|
g_loss = -torch.mean(logits_fake)
|
||||||
|
if self.training:
|
||||||
|
d_weight = self.calculate_adaptive_weight(
|
||||||
|
nll_loss, g_loss, last_layer=last_layer
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
d_weight = torch.tensor(1.0)
|
||||||
|
else:
|
||||||
|
d_weight = torch.tensor(0.0)
|
||||||
|
g_loss = torch.tensor(0.0, requires_grad=True)
|
||||||
|
|
||||||
|
loss = weighted_nll_loss + d_weight * self.disc_factor * g_loss
|
||||||
|
log = dict()
|
||||||
|
for k in regularization_log:
|
||||||
|
if k in self.regularization_weights:
|
||||||
|
loss = loss + self.regularization_weights[k] * regularization_log[k]
|
||||||
|
if k in self.additional_log_keys:
|
||||||
|
log[f"{split}/{k}"] = regularization_log[k].detach().float().mean()
|
||||||
|
|
||||||
|
log.update(
|
||||||
|
{
|
||||||
|
f"{split}/loss/total": loss.clone().detach().mean(),
|
||||||
|
f"{split}/loss/nll": nll_loss.detach().mean(),
|
||||||
|
f"{split}/loss/rec": rec_loss.detach().mean(),
|
||||||
|
f"{split}/loss/g": g_loss.detach().mean(),
|
||||||
|
f"{split}/scalars/logvar": self.logvar.detach(),
|
||||||
|
f"{split}/scalars/d_weight": d_weight.detach(),
|
||||||
|
}
|
||||||
|
)
|
||||||
|
|
||||||
|
return loss, log
|
||||||
|
elif optimizer_idx == 1:
|
||||||
|
# second pass for discriminator update
|
||||||
|
logits_real = self.discriminator(inputs.contiguous().detach())
|
||||||
|
logits_fake = self.discriminator(reconstructions.contiguous().detach())
|
||||||
|
|
||||||
|
if global_step >= self.discriminator_iter_start or not self.training:
|
||||||
|
d_loss = self.disc_factor * self.disc_loss(logits_real, logits_fake)
|
||||||
|
else:
|
||||||
|
d_loss = torch.tensor(0.0, requires_grad=True)
|
||||||
|
|
||||||
|
log = {
|
||||||
|
f"{split}/loss/disc": d_loss.clone().detach().mean(),
|
||||||
|
f"{split}/logits/real": logits_real.detach().mean(),
|
||||||
|
f"{split}/logits/fake": logits_fake.detach().mean(),
|
||||||
|
}
|
||||||
|
return d_loss, log
|
||||||
|
else:
|
||||||
|
raise NotImplementedError(f"Unknown optimizer_idx {optimizer_idx}")
|
||||||
|
|
||||||
|
def get_nll_loss(
|
||||||
|
self,
|
||||||
|
rec_loss: torch.Tensor,
|
||||||
|
weights: Optional[Union[float, torch.Tensor]] = None,
|
||||||
|
) -> Tuple[torch.Tensor, torch.Tensor]:
|
||||||
|
nll_loss = rec_loss / torch.exp(self.logvar) + self.logvar
|
||||||
|
weighted_nll_loss = nll_loss
|
||||||
|
if weights is not None:
|
||||||
|
weighted_nll_loss = weights * nll_loss
|
||||||
|
weighted_nll_loss = torch.sum(weighted_nll_loss) / weighted_nll_loss.shape[0]
|
||||||
|
nll_loss = torch.sum(nll_loss) / nll_loss.shape[0]
|
||||||
|
|
||||||
|
return nll_loss, weighted_nll_loss
|
||||||
@@ -0,0 +1,73 @@
|
|||||||
|
import torch
|
||||||
|
import torch.nn as nn
|
||||||
|
|
||||||
|
from ....util import default, instantiate_from_config
|
||||||
|
from ..lpips.loss.lpips import LPIPS
|
||||||
|
|
||||||
|
|
||||||
|
class LatentLPIPS(nn.Module):
|
||||||
|
def __init__(
|
||||||
|
self,
|
||||||
|
decoder_config,
|
||||||
|
perceptual_weight=1.0,
|
||||||
|
latent_weight=1.0,
|
||||||
|
scale_input_to_tgt_size=False,
|
||||||
|
scale_tgt_to_input_size=False,
|
||||||
|
perceptual_weight_on_inputs=0.0,
|
||||||
|
):
|
||||||
|
super().__init__()
|
||||||
|
self.scale_input_to_tgt_size = scale_input_to_tgt_size
|
||||||
|
self.scale_tgt_to_input_size = scale_tgt_to_input_size
|
||||||
|
self.init_decoder(decoder_config)
|
||||||
|
self.perceptual_loss = LPIPS().eval()
|
||||||
|
self.perceptual_weight = perceptual_weight
|
||||||
|
self.latent_weight = latent_weight
|
||||||
|
self.perceptual_weight_on_inputs = perceptual_weight_on_inputs
|
||||||
|
|
||||||
|
def init_decoder(self, config):
|
||||||
|
self.decoder = instantiate_from_config(config)
|
||||||
|
if hasattr(self.decoder, "encoder"):
|
||||||
|
del self.decoder.encoder
|
||||||
|
|
||||||
|
def forward(self, latent_inputs, latent_predictions, image_inputs, split="train"):
|
||||||
|
log = dict()
|
||||||
|
loss = (latent_inputs - latent_predictions) ** 2
|
||||||
|
log[f"{split}/latent_l2_loss"] = loss.mean().detach()
|
||||||
|
image_reconstructions = None
|
||||||
|
if self.perceptual_weight > 0.0:
|
||||||
|
image_reconstructions = self.decoder.decode(latent_predictions)
|
||||||
|
image_targets = self.decoder.decode(latent_inputs)
|
||||||
|
perceptual_loss = self.perceptual_loss(
|
||||||
|
image_targets.contiguous(), image_reconstructions.contiguous()
|
||||||
|
)
|
||||||
|
loss = (
|
||||||
|
self.latent_weight * loss.mean()
|
||||||
|
+ self.perceptual_weight * perceptual_loss.mean()
|
||||||
|
)
|
||||||
|
log[f"{split}/perceptual_loss"] = perceptual_loss.mean().detach()
|
||||||
|
|
||||||
|
if self.perceptual_weight_on_inputs > 0.0:
|
||||||
|
image_reconstructions = default(
|
||||||
|
image_reconstructions, self.decoder.decode(latent_predictions)
|
||||||
|
)
|
||||||
|
if self.scale_input_to_tgt_size:
|
||||||
|
image_inputs = torch.nn.functional.interpolate(
|
||||||
|
image_inputs,
|
||||||
|
image_reconstructions.shape[2:],
|
||||||
|
mode="bicubic",
|
||||||
|
antialias=True,
|
||||||
|
)
|
||||||
|
elif self.scale_tgt_to_input_size:
|
||||||
|
image_reconstructions = torch.nn.functional.interpolate(
|
||||||
|
image_reconstructions,
|
||||||
|
image_inputs.shape[2:],
|
||||||
|
mode="bicubic",
|
||||||
|
antialias=True,
|
||||||
|
)
|
||||||
|
|
||||||
|
perceptual_loss2 = self.perceptual_loss(
|
||||||
|
image_inputs.contiguous(), image_reconstructions.contiguous()
|
||||||
|
)
|
||||||
|
loss = loss + self.perceptual_weight_on_inputs * perceptual_loss2.mean()
|
||||||
|
log[f"{split}/perceptual_loss_on_inputs"] = perceptual_loss2.mean().detach()
|
||||||
|
return loss, log
|
||||||
@@ -0,0 +1 @@
|
|||||||
|
vgg.pth
|
||||||
@@ -0,0 +1,23 @@
|
|||||||
|
Copyright (c) 2018, Richard Zhang, Phillip Isola, Alexei A. Efros, Eli Shechtman, Oliver Wang
|
||||||
|
All rights reserved.
|
||||||
|
|
||||||
|
Redistribution and use in source and binary forms, with or without
|
||||||
|
modification, are permitted provided that the following conditions are met:
|
||||||
|
|
||||||
|
* Redistributions of source code must retain the above copyright notice, this
|
||||||
|
list of conditions and the following disclaimer.
|
||||||
|
|
||||||
|
* Redistributions in binary form must reproduce the above copyright notice,
|
||||||
|
this list of conditions and the following disclaimer in the documentation
|
||||||
|
and/or other materials provided with the distribution.
|
||||||
|
|
||||||
|
THIS SOFTWARE IS PROVIDED BY THE COPYRIGHT HOLDERS AND CONTRIBUTORS "AS IS"
|
||||||
|
AND ANY EXPRESS OR IMPLIED WARRANTIES, INCLUDING, BUT NOT LIMITED TO, THE
|
||||||
|
IMPLIED WARRANTIES OF MERCHANTABILITY AND FITNESS FOR A PARTICULAR PURPOSE ARE
|
||||||
|
DISCLAIMED. IN NO EVENT SHALL THE COPYRIGHT HOLDER OR CONTRIBUTORS BE LIABLE
|
||||||
|
FOR ANY DIRECT, INDIRECT, INCIDENTAL, SPECIAL, EXEMPLARY, OR CONSEQUENTIAL
|
||||||
|
DAMAGES (INCLUDING, BUT NOT LIMITED TO, PROCUREMENT OF SUBSTITUTE GOODS OR
|
||||||
|
SERVICES; LOSS OF USE, DATA, OR PROFITS; OR BUSINESS INTERRUPTION) HOWEVER
|
||||||
|
CAUSED AND ON ANY THEORY OF LIABILITY, WHETHER IN CONTRACT, STRICT LIABILITY,
|
||||||
|
OR TORT (INCLUDING NEGLIGENCE OR OTHERWISE) ARISING IN ANY WAY OUT OF THE USE
|
||||||
|
OF THIS SOFTWARE, EVEN IF ADVISED OF THE POSSIBILITY OF SUCH DAMAGE.
|
||||||
@@ -0,0 +1,147 @@
|
|||||||
|
"""Stripped version of https://github.com/richzhang/PerceptualSimilarity/tree/master/models"""
|
||||||
|
|
||||||
|
from collections import namedtuple
|
||||||
|
|
||||||
|
import torch
|
||||||
|
import torch.nn as nn
|
||||||
|
from torchvision import models
|
||||||
|
|
||||||
|
from ..util import get_ckpt_path
|
||||||
|
|
||||||
|
|
||||||
|
class LPIPS(nn.Module):
|
||||||
|
# Learned perceptual metric
|
||||||
|
def __init__(self, use_dropout=True):
|
||||||
|
super().__init__()
|
||||||
|
self.scaling_layer = ScalingLayer()
|
||||||
|
self.chns = [64, 128, 256, 512, 512] # vg16 features
|
||||||
|
self.net = vgg16(pretrained=True, requires_grad=False)
|
||||||
|
self.lin0 = NetLinLayer(self.chns[0], use_dropout=use_dropout)
|
||||||
|
self.lin1 = NetLinLayer(self.chns[1], use_dropout=use_dropout)
|
||||||
|
self.lin2 = NetLinLayer(self.chns[2], use_dropout=use_dropout)
|
||||||
|
self.lin3 = NetLinLayer(self.chns[3], use_dropout=use_dropout)
|
||||||
|
self.lin4 = NetLinLayer(self.chns[4], use_dropout=use_dropout)
|
||||||
|
self.load_from_pretrained()
|
||||||
|
for param in self.parameters():
|
||||||
|
param.requires_grad = False
|
||||||
|
|
||||||
|
def load_from_pretrained(self, name="vgg_lpips"):
|
||||||
|
ckpt = get_ckpt_path(name, ".sgm/modules/autoencoding/lpips/loss")
|
||||||
|
self.load_state_dict(
|
||||||
|
torch.load(ckpt, map_location=torch.device("cpu")), strict=False
|
||||||
|
)
|
||||||
|
print("loaded pretrained LPIPS loss from {}".format(ckpt))
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def from_pretrained(cls, name="vgg_lpips"):
|
||||||
|
if name != "vgg_lpips":
|
||||||
|
raise NotImplementedError
|
||||||
|
model = cls()
|
||||||
|
ckpt = get_ckpt_path(name)
|
||||||
|
model.load_state_dict(
|
||||||
|
torch.load(ckpt, map_location=torch.device("cpu")), strict=False
|
||||||
|
)
|
||||||
|
return model
|
||||||
|
|
||||||
|
def forward(self, input, target):
|
||||||
|
in0_input, in1_input = (self.scaling_layer(input), self.scaling_layer(target))
|
||||||
|
outs0, outs1 = self.net(in0_input), self.net(in1_input)
|
||||||
|
feats0, feats1, diffs = {}, {}, {}
|
||||||
|
lins = [self.lin0, self.lin1, self.lin2, self.lin3, self.lin4]
|
||||||
|
for kk in range(len(self.chns)):
|
||||||
|
feats0[kk], feats1[kk] = normalize_tensor(outs0[kk]), normalize_tensor(
|
||||||
|
outs1[kk]
|
||||||
|
)
|
||||||
|
diffs[kk] = (feats0[kk] - feats1[kk]) ** 2
|
||||||
|
|
||||||
|
res = [
|
||||||
|
spatial_average(lins[kk].model(diffs[kk]), keepdim=True)
|
||||||
|
for kk in range(len(self.chns))
|
||||||
|
]
|
||||||
|
val = res[0]
|
||||||
|
for l in range(1, len(self.chns)):
|
||||||
|
val += res[l]
|
||||||
|
return val
|
||||||
|
|
||||||
|
|
||||||
|
class ScalingLayer(nn.Module):
|
||||||
|
def __init__(self):
|
||||||
|
super(ScalingLayer, self).__init__()
|
||||||
|
self.register_buffer(
|
||||||
|
"shift", torch.Tensor([-0.030, -0.088, -0.188])[None, :, None, None]
|
||||||
|
)
|
||||||
|
self.register_buffer(
|
||||||
|
"scale", torch.Tensor([0.458, 0.448, 0.450])[None, :, None, None]
|
||||||
|
)
|
||||||
|
|
||||||
|
def forward(self, inp):
|
||||||
|
return (inp - self.shift) / self.scale
|
||||||
|
|
||||||
|
|
||||||
|
class NetLinLayer(nn.Module):
|
||||||
|
"""A single linear layer which does a 1x1 conv"""
|
||||||
|
|
||||||
|
def __init__(self, chn_in, chn_out=1, use_dropout=False):
|
||||||
|
super(NetLinLayer, self).__init__()
|
||||||
|
layers = (
|
||||||
|
[
|
||||||
|
nn.Dropout(),
|
||||||
|
]
|
||||||
|
if (use_dropout)
|
||||||
|
else []
|
||||||
|
)
|
||||||
|
layers += [
|
||||||
|
nn.Conv2d(chn_in, chn_out, 1, stride=1, padding=0, bias=False),
|
||||||
|
]
|
||||||
|
self.model = nn.Sequential(*layers)
|
||||||
|
|
||||||
|
|
||||||
|
class vgg16(torch.nn.Module):
|
||||||
|
def __init__(self, requires_grad=False, pretrained=True):
|
||||||
|
super(vgg16, self).__init__()
|
||||||
|
vgg_pretrained_features = models.vgg16(pretrained=pretrained).features
|
||||||
|
self.slice1 = torch.nn.Sequential()
|
||||||
|
self.slice2 = torch.nn.Sequential()
|
||||||
|
self.slice3 = torch.nn.Sequential()
|
||||||
|
self.slice4 = torch.nn.Sequential()
|
||||||
|
self.slice5 = torch.nn.Sequential()
|
||||||
|
self.N_slices = 5
|
||||||
|
for x in range(4):
|
||||||
|
self.slice1.add_module(str(x), vgg_pretrained_features[x])
|
||||||
|
for x in range(4, 9):
|
||||||
|
self.slice2.add_module(str(x), vgg_pretrained_features[x])
|
||||||
|
for x in range(9, 16):
|
||||||
|
self.slice3.add_module(str(x), vgg_pretrained_features[x])
|
||||||
|
for x in range(16, 23):
|
||||||
|
self.slice4.add_module(str(x), vgg_pretrained_features[x])
|
||||||
|
for x in range(23, 30):
|
||||||
|
self.slice5.add_module(str(x), vgg_pretrained_features[x])
|
||||||
|
if not requires_grad:
|
||||||
|
for param in self.parameters():
|
||||||
|
param.requires_grad = False
|
||||||
|
|
||||||
|
def forward(self, X):
|
||||||
|
h = self.slice1(X)
|
||||||
|
h_relu1_2 = h
|
||||||
|
h = self.slice2(h)
|
||||||
|
h_relu2_2 = h
|
||||||
|
h = self.slice3(h)
|
||||||
|
h_relu3_3 = h
|
||||||
|
h = self.slice4(h)
|
||||||
|
h_relu4_3 = h
|
||||||
|
h = self.slice5(h)
|
||||||
|
h_relu5_3 = h
|
||||||
|
vgg_outputs = namedtuple(
|
||||||
|
"VggOutputs", ["relu1_2", "relu2_2", "relu3_3", "relu4_3", "relu5_3"]
|
||||||
|
)
|
||||||
|
out = vgg_outputs(h_relu1_2, h_relu2_2, h_relu3_3, h_relu4_3, h_relu5_3)
|
||||||
|
return out
|
||||||
|
|
||||||
|
|
||||||
|
def normalize_tensor(x, eps=1e-10):
|
||||||
|
norm_factor = torch.sqrt(torch.sum(x**2, dim=1, keepdim=True))
|
||||||
|
return x / (norm_factor + eps)
|
||||||
|
|
||||||
|
|
||||||
|
def spatial_average(x, keepdim=True):
|
||||||
|
return x.mean([2, 3], keepdim=keepdim)
|
||||||