initial commit

This commit is contained in:
kijai
2024-09-30 00:29:02 +03:00
parent 2cf0731c6e
commit 3d428a5ef2
132 changed files with 11850 additions and 0 deletions
+146
View File
@@ -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
+7
View File
@@ -0,0 +1,7 @@
# ComfyUI wrapper nodes for LVCD:
Original repo:
https://github.com/luckyhzt/LVCD
+3
View File
@@ -0,0 +1,3 @@
from .nodes import NODE_CLASS_MAPPINGS, NODE_DISPLAY_NAME_MAPPINGS
__all__ = ["NODE_CLASS_MAPPINGS", "NODE_DISPLAY_NAME_MAPPINGS"]
+250
View File
@@ -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
+197
View File
@@ -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)
+472
View File
@@ -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
Binary file not shown.

After

Width:  |  Height:  |  Size: 103 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 103 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 100 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 98 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 98 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 99 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 99 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 101 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 101 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 102 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 101 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 103 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 104 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 104 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 100 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 99 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 98 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 98 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 98 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 98 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 98 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 98 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 99 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 99 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 99 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 99 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 100 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 99 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 99 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 102 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 102 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 105 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 107 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 106 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 105 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 104 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 105 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 106 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 106 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 106 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 104 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 104 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 104 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 104 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 108 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 107 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 108 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 109 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 111 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 107 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 106 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 107 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 107 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 107 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 107 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 107 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 107 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 107 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 107 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 108 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 107 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 107 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 106 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 106 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 105 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 105 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 105 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 106 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 106 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 106 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 106 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 106 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 107 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 107 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 107 KiB

+1054
View File
File diff suppressed because it is too large Load Diff
+301
View File
@@ -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
+209
View File
@@ -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",
}
+4
View File
@@ -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
+4
View File
@@ -0,0 +1,4 @@
from .models import AutoencodingEngine, DiffusionEngine
from .util import get_configs_path, instantiate_from_config
__version__ = "0.1.0"
+2
View File
@@ -0,0 +1,2 @@
from .autoencoder import AutoencodingEngine
from .diffusion import DiffusionEngine
+615
View File
@@ -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,
)
+356
View File
@@ -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
+6
View File
@@ -0,0 +1,6 @@
from .encoders.modules import GeneralConditioner
UNCONDITIONAL_CONFIG = {
"target": ".sgm.modules.GeneralConditioner",
"params": {"emb_models": []},
}
File diff suppressed because it is too large Load Diff
@@ -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
+73
View File
@@ -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)

Some files were not shown because too many files have changed in this diff Show More