Initial commit

This commit is contained in:
Kijai
2024-04-11 14:16:38 +03:00
parent 67d2eddb23
commit 73bfb99844
10 changed files with 1849 additions and 0 deletions
+9
View File
@@ -0,0 +1,9 @@
pretrained_models/
example_data/
results/
*.zip
.vscode/
.hypothesis/
*.pt
__pycache__
*.pyc
+21
View File
@@ -0,0 +1,21 @@
MIT License
Copyright (c) Microsoft Corporation.
Permission is hereby granted, free of charge, to any person obtaining a copy
of this software and associated documentation files (the "Software"), to deal
in the Software without restriction, including without limitation the rights
to use, copy, modify, merge, publish, distribute, sublicense, and/or sell
copies of the Software, and to permit persons to whom the Software is
furnished to do so, subject to the following conditions:
The above copyright notice and this permission notice shall be included in all
copies or substantial portions of the Software.
THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR
IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY,
FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE
AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER
LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM,
OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE
SOFTWARE
+16
View File
@@ -0,0 +1,16 @@
# ComfyUI wrapper node to test LaVi-Bridge using Diffusers
# Installing
Either use the Manager and it's install from git -feature, or clone this repo to custom_nodes and run:
`pip install -r requirements.txt`
or if you use portable (run this in ComfyUI_windows_portable -folder):
`python_embeded\python.exe -m pip install -r ComfyUI\custom_nodes\ComfyUI-Lavi-Bridge-Wrapper\requirements.txt`
The following is autodownloaded:
https://huggingface.co/Kijai/t5-large-encoder-only-bf16/ to `ComfyUI/models/t5_model/``
https://huggingface.co/shihaozhao/LaVi-Bridge/ to `ComfyUI/models/lavibridge`
+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"]
+70
View File
@@ -0,0 +1,70 @@
model:
base_learning_rate: 1.0e-04
target: ldm.models.diffusion.ddpm.LatentDiffusion
params:
linear_start: 0.00085
linear_end: 0.0120
num_timesteps_cond: 1
log_every_t: 200
timesteps: 1000
first_stage_key: "jpg"
cond_stage_key: "txt"
image_size: 64
channels: 4
cond_stage_trainable: false # Note: different from the one we trained before
conditioning_key: crossattn
monitor: val/loss_simple_ema
scale_factor: 0.18215
use_ema: False
scheduler_config: # 10000 warmup steps
target: ldm.lr_scheduler.LambdaLinearScheduler
params:
warm_up_steps: [ 10000 ]
cycle_lengths: [ 10000000000000 ] # incredibly large number to prevent corner cases
f_start: [ 1.e-6 ]
f_max: [ 1. ]
f_min: [ 1. ]
unet_config:
target: ldm.modules.diffusionmodules.openaimodel.UNetModel
params:
image_size: 32 # unused
in_channels: 4
out_channels: 4
model_channels: 320
attention_resolutions: [ 4, 2, 1 ]
num_res_blocks: 2
channel_mult: [ 1, 2, 4, 4 ]
num_heads: 8
use_spatial_transformer: True
transformer_depth: 1
context_dim: 768
use_checkpoint: True
legacy: False
first_stage_config:
target: ldm.models.autoencoder.AutoencoderKL
params:
embed_dim: 4
monitor: val/rec_loss
ddconfig:
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
cond_stage_config:
target: ldm.modules.encoders.modules.FrozenCLIPEmbedder
+249
View File
@@ -0,0 +1,249 @@
{
"last_node_id": 5,
"last_link_id": 5,
"nodes": [
{
"id": 3,
"type": "CheckpointLoaderSimple",
"pos": [
373,
314
],
"size": [
320.97000488281265,
98
],
"flags": {},
"order": 0,
"mode": 0,
"outputs": [
{
"name": "MODEL",
"type": "MODEL",
"links": [
2
],
"shape": 3
},
{
"name": "CLIP",
"type": "CLIP",
"links": null,
"shape": 3
},
{
"name": "VAE",
"type": "VAE",
"links": [
3
],
"shape": 3,
"slot_index": 2
}
],
"properties": {
"Node name for S&R": "CheckpointLoaderSimple"
},
"widgets_values": [
"1_5/dreamshaper_8.safetensors"
]
},
{
"id": 2,
"type": "lavibridge_model_loader",
"pos": [
763,
322
],
"size": {
"0": 210,
"1": 46
},
"flags": {},
"order": 2,
"mode": 0,
"inputs": [
{
"name": "model",
"type": "MODEL",
"link": 2,
"slot_index": 0
},
{
"name": "vae",
"type": "VAE",
"link": 3
}
],
"outputs": [
{
"name": "lavibridge",
"type": "LAVIBRIDGE",
"links": [
1
],
"shape": 3
}
],
"properties": {
"Node name for S&R": "lavibridge_model_loader"
}
},
{
"id": 4,
"type": "PreviewImage",
"pos": [
1388,
321
],
"size": [
558.763644131747,
569.452726537531
],
"flags": {},
"order": 4,
"mode": 0,
"inputs": [
{
"name": "images",
"type": "IMAGE",
"link": 4
}
],
"properties": {
"Node name for S&R": "PreviewImage"
}
},
{
"id": 5,
"type": "lavi_bridge_t5_encoder",
"pos": [
376,
484
],
"size": [
322.0363714044744,
200.2709083557129
],
"flags": {},
"order": 1,
"mode": 0,
"outputs": [
{
"name": "t5_embeds",
"type": "T5EMBEDS",
"links": [
5
],
"shape": 3
}
],
"properties": {
"Node name for S&R": "lavi_bridge_t5_encoder"
},
"widgets_values": [
"Oppenheimer sits on the beach on a chair, watching a nuclear exposition with a huge mushroom cloud, 120mm, best quality, masterpiece, extremely detailed, 4k resolution",
77
]
},
{
"id": 1,
"type": "lavibridge_sampler",
"pos": [
1021,
320
],
"size": {
"0": 315,
"1": 246
},
"flags": {},
"order": 3,
"mode": 0,
"inputs": [
{
"name": "lavibridge_model",
"type": "LAVIBRIDGE",
"link": 1,
"slot_index": 0
},
{
"name": "t5_embeds",
"type": "T5EMBEDS",
"link": 5,
"slot_index": 1
}
],
"outputs": [
{
"name": "images",
"type": "IMAGE",
"links": [
4
],
"shape": 3,
"slot_index": 0
}
],
"properties": {
"Node name for S&R": "lavibridge_sampler"
},
"widgets_values": [
512,
512,
4,
25,
7.5,
0,
"fixed",
"UniPCMultistepScheduler"
]
}
],
"links": [
[
1,
2,
0,
1,
0,
"LAVIBRIDGE"
],
[
2,
3,
0,
2,
0,
"MODEL"
],
[
3,
3,
2,
2,
1,
"VAE"
],
[
4,
1,
0,
4,
0,
"IMAGE"
],
[
5,
5,
0,
1,
1,
"T5EMBEDS"
]
],
"groups": [],
"config": {},
"extra": {},
"version": 0.4
}
+62
View File
@@ -0,0 +1,62 @@
from dataclasses import dataclass
from inspect import isfunction
import torch
import torch.nn as nn
import torch.nn.functional as F
from diffusers.utils import BaseOutput
from diffusers.models.modeling_utils import ModelMixin
from diffusers.configuration_utils import ConfigMixin, register_to_config
def default(val, d):
if val is not None: return val
return d() if isfunction(d) else d
class GEGLU(nn.Module):
def __init__(self, dim_in, dim_out):
super().__init__()
self.proj = nn.Linear(dim_in, dim_out * 2)
def forward(self, x):
x, gate = self.proj(x).chunk(2, dim=-1)
return x * F.gelu(gate)
class FeedForward(nn.Module):
def __init__(self, dim, dim_out, mult=4, dropout=0.1):
super().__init__()
inner_dim = int(dim * mult)
dim_out = default(dim_out, dim)
project_in = GEGLU(dim, inner_dim)
self.net = nn.Sequential(
project_in,
nn.Dropout(dropout),
nn.Linear(inner_dim, dim_out)
)
def forward(self, x):
return self.net(x)
@dataclass
class TextAdapterOutput(BaseOutput):
sample: torch.FloatTensor
class TextAdapter(ModelMixin, ConfigMixin):
@register_to_config
def __init__(self, in_dim, int_dim, out_dim):
super().__init__()
self.in_dim = in_dim
self.ff1 = FeedForward(in_dim, int_dim)
self.ff2 = FeedForward(int_dim, out_dim)
self.norm1 = nn.LayerNorm(in_dim)
self.norm2 = nn.LayerNorm(int_dim)
def forward(self, x):
x = self.ff1(self.norm1(x))
x = self.ff2(self.norm2(x))
return TextAdapterOutput(x)
+1106
View File
File diff suppressed because it is too large Load Diff
+310
View File
@@ -0,0 +1,310 @@
import os
from tqdm.auto import tqdm
try:
from diffusers import (
DDIMScheduler,
DPMSolverMultistepScheduler,
EulerDiscreteScheduler,
EulerAncestralDiscreteScheduler,
AutoencoderKL,
LCMScheduler,
DDPMScheduler,
DEISMultistepScheduler,
PNDMScheduler,
UniPCMultistepScheduler
)
from diffusers.loaders.single_file_utils import (
convert_ldm_vae_checkpoint,
convert_ldm_unet_checkpoint,
create_vae_diffusers_config,
create_unet_diffusers_config
)
except:
print("Diffusers version too old. Please update to 0.26.0 minimum.")
import torch
from contextlib import nullcontext
from diffusers import AutoencoderKL, UNet2DConditionModel
from transformers import AutoTokenizer, T5EncoderModel
from omegaconf import OmegaConf
from .modules.lora import monkeypatch_or_replace_lora_extended
from .modules.adapters import TextAdapter
import folder_paths
import comfy.latent_formats
import comfy.model_management as mm
script_directory = os.path.dirname(os.path.abspath(__file__))
class lavibridge_model_loader:
@classmethod
def INPUT_TYPES(s):
return {"required": {
"model": ("MODEL",),
"vae": ("VAE",),
},
}
RETURN_TYPES = ("LAVIBRIDGE",)
RETURN_NAMES = ("lavibridge",)
FUNCTION = "loadmodel"
CATEGORY = "LaVI-BridgeWrapper"
def loadmodel(self, model, vae):
mm.soft_empty_cache()
dtype = mm.unet_dtype()
vae_dtype = mm.vae_dtype()
custom_config = {
'model': model,
'vae': vae,
}
if not hasattr(self, 'model') or self.model == None or custom_config != self.current_config:
pbar = comfy.utils.ProgressBar(5)
self.current_config = custom_config
# config paths
original_config = OmegaConf.load(os.path.join(script_directory, f"configs/v1-inference.yaml"))
# load models
lavibridge_folder = os.path.join(folder_paths.models_dir,'lavibridge')
lora_vis_path = os.path.join(lavibridge_folder, 't5_unet', 'lora_vis.pt')
if not os.path.exists(lora_vis_path):
print(f"Downloading LaVi-Bridge from https://huggingface.co/shihaozhao/LaVi-Bridge {lavibridge_folder}")
from huggingface_hub import snapshot_download
snapshot_download(repo_id="shihaozhao/LaVi-Bridge", allow_patterns=["*t5_unet*"],local_dir=lavibridge_folder, local_dir_use_symlinks=False)
pbar.update(1)
# get state dict from comfy models
load_models = [model]
comfy.model_management.load_models_gpu(load_models)
sd = model.model.state_dict_for_saving(None, vae.get_sd(), None)
pbar.update(1)
# 1. vae
converted_vae_config = create_vae_diffusers_config(original_config, image_size=512)
converted_vae = convert_ldm_vae_checkpoint(sd, converted_vae_config)
vae = AutoencoderKL(**converted_vae_config)
vae.load_state_dict(converted_vae, strict=False)
vae.to(vae_dtype).eval()
pbar.update(1)
# 2. unet
converted_unet_config = create_unet_diffusers_config(original_config, image_size=512)
converted_unet = convert_ldm_unet_checkpoint(sd, converted_unet_config)
unet = UNet2DConditionModel(**converted_unet_config)
unet.load_state_dict(converted_unet, strict=False)
unet.eval()
pbar.update(1)
# LoRA
monkeypatch_or_replace_lora_extended(
unet,
torch.load(lora_vis_path),
r=32,
target_replace_module={"ResnetBlock2D", "CrossAttention", "Attention", "GEGLU"},
)
unet.to(dtype)
pbar.update(1)
lavibridge_model = {
'unet': unet,
'vae': vae,
}
return (lavibridge_model,)
class lavi_bridge_t5_encoder:
@classmethod
def INPUT_TYPES(s):
return {"required": {
"prompt": ("STRING", {"multiline": True, "default": "A vivid red book with a smooth, matte cover lies next to a glossy yellow vase. The vase, with a slightly curved silhouette, stands on a dark wood table with a noticeable grain pattern. The book appears slightly worn at the edges, suggesting frequent use, while the vase holds a fresh array of multicolored wildflowers.",}),
"max_length": ("INT", {"default": 77, "min": 1, "max": 512, "step": 1}),
},
}
RETURN_TYPES = ("T5EMBEDS",)
RETURN_NAMES = ("t5_embeds",)
FUNCTION = "process"
CATEGORY = "LaVI-BridgeWrapper"
def process(self, prompt, max_length):
device = mm.get_torch_device()
offload_device = mm.unet_offload_device()
mm.soft_empty_cache()
#dtype = mm.unet_dtype()
dtype = torch.bfloat16
if not hasattr(self, "text_encoder"):
#t5
t5_path = os.path.join(folder_paths.models_dir,'t5_model', 't5-large-encoder-only-bf16')
if not os.path.exists(t5_path):
from huggingface_hub import snapshot_download
snapshot_download(repo_id="Kijai/t5-large-encoder-only-bf16", local_dir=t5_path, local_dir_use_symlinks=False)
#adapter
adapter_folder = os.path.join(folder_paths.models_dir,'lavibridge')
adapter_path = os.path.join(adapter_folder, 't5_unet','adapter')
if not os.path.exists(adapter_path):
print(f"Downloading LaVi-Bridge from https://huggingface.co/shihaozhao/LaVi-Bridge {adapter_folder}")
from huggingface_hub import snapshot_download
snapshot_download(repo_id="shihaozhao/LaVi-Bridge", allow_patterns=["t5_unet"],local_dir=adapter_folder, local_dir_use_symlinks=False)
lora_text_path = os.path.join(adapter_folder, 't5_unet', 'lora_text.pt')
self.adapter = TextAdapter.from_pretrained(adapter_path).eval().to(dtype)
self.tokenizer = AutoTokenizer.from_pretrained(t5_path)
self.text_encoder = T5EncoderModel.from_pretrained(t5_path).eval().to(dtype)
monkeypatch_or_replace_lora_extended(
self.text_encoder,
torch.load(lora_text_path),
r=32,
target_replace_module = {"T5Attention"},
)
self.adapter.to(device)
self.text_encoder.to(device)
autocast_condition = (dtype != torch.float32) and not mm.is_device_mps(device)
with torch.autocast(mm.get_autocast_device(device), dtype=dtype) if autocast_condition else nullcontext():
text_ids = self.tokenizer(prompt, padding="max_length", max_length=max_length, return_tensors="pt", truncation=True).input_ids.to(device)
text_embeddings = self.text_encoder(input_ids=text_ids)[0]
text_embeddings = self.adapter(text_embeddings).sample
uncond_input = self.tokenizer([""], padding="max_length", max_length=max_length, return_tensors="pt")
uncond_embeddings = self.text_encoder(uncond_input.input_ids.to(device))[0]
uncond_embeddings = self.adapter(uncond_embeddings).sample
text_embeddings = torch.cat([uncond_embeddings, text_embeddings])
self.adapter.to(offload_device)
self.text_encoder.to(offload_device)
return (text_embeddings,)
class lavibridge_sampler:
@classmethod
def INPUT_TYPES(s):
return {"required": {
"lavibridge_model": ("LAVIBRIDGE",),
"t5_embeds": ("T5EMBEDS",),
"width": ("INT", {"default": 512, "min": 64, "max": 2048, "step": 64}),
"height": ("INT", {"default": 512, "min": 64, "max": 2048, "step": 64}),
"batch_size": ("INT", {"default": 1, "min": 1, "max": 256, "step": 1}),
"steps": ("INT", {"default": 25, "min": 1, "max": 200, "step": 1}),
"guidance_scale": ("FLOAT", {"default": 7.5, "min": 0.0, "max": 20.0, "step": 0.01}),
"seed": ("INT", {"default": 0, "min": 0, "max": 0xffffffffffffffff}),
"scheduler": (
[
'DPMSolverMultistepScheduler',
'DPMSolverMultistepScheduler_SDE_karras',
'DDPMScheduler',
'LCMScheduler',
'PNDMScheduler',
'DEISMultistepScheduler',
'EulerDiscreteScheduler',
'EulerAncestralDiscreteScheduler',
'UniPCMultistepScheduler',
'DDIMScheduler',
], {
"default": 'DPMSolverMultistepScheduler'
}),
},
}
RETURN_TYPES = ("IMAGE",)
RETURN_NAMES = ("images",)
FUNCTION = "process"
CATEGORY = "LaVI-BridgeWrapper"
def process(self, lavibridge_model, t5_embeds, width, height, batch_size, steps, guidance_scale, seed, scheduler):
device = mm.get_torch_device()
offload_device = mm.unet_offload_device()
mm.unload_all_models()
mm.soft_empty_cache()
torch.manual_seed(seed)
dtype = mm.unet_dtype()
unet = lavibridge_model["unet"]
vae = lavibridge_model["vae"]
scheduler_config = {
'num_train_timesteps': 1000,
'beta_start': 0.00085,
'beta_end': 0.012,
'beta_schedule': "scaled_linear",
'steps_offset': 1,
}
if scheduler == 'DPMSolverMultistepScheduler':
noise_scheduler = DPMSolverMultistepScheduler(**scheduler_config)
elif scheduler == 'DPMSolverMultistepScheduler_SDE_karras':
scheduler_config.update({"algorithm_type": "sde-dpmsolver++"})
scheduler_config.update({"use_karras_sigmas": True})
noise_scheduler = DPMSolverMultistepScheduler(**scheduler_config)
elif scheduler == 'DDPMScheduler':
noise_scheduler = DDPMScheduler(**scheduler_config)
elif scheduler == 'LCMScheduler':
noise_scheduler = LCMScheduler(**scheduler_config)
elif scheduler == 'PNDMScheduler':
scheduler_config.update({"set_alpha_to_one": False})
scheduler_config.update({"trained_betas": None})
noise_scheduler = PNDMScheduler(**scheduler_config)
elif scheduler == 'DEISMultistepScheduler':
noise_scheduler = DEISMultistepScheduler(**scheduler_config)
elif scheduler == 'EulerDiscreteScheduler':
noise_scheduler = EulerDiscreteScheduler(**scheduler_config)
elif scheduler == 'EulerAncestralDiscreteScheduler':
noise_scheduler = EulerAncestralDiscreteScheduler(**scheduler_config)
elif scheduler == 'UniPCMultistepScheduler':
noise_scheduler = UniPCMultistepScheduler(**scheduler_config)
elif scheduler == 'DDIMScheduler':
noise_scheduler = DDIMScheduler(**scheduler_config)
unet.to(device)
autocast_condition = (dtype != torch.float32) and not mm.is_device_mps(device)
with torch.autocast(mm.get_autocast_device(device), dtype=dtype) if autocast_condition else nullcontext():
# Latent preparation
vae.to(device)
latents = torch.randn((batch_size, unet.in_channels, height // 8, width // 8)).to(device)
latents = latents * noise_scheduler.init_noise_sigma
vae.to(offload_device)
t5_embeds_repeated = t5_embeds.repeat_interleave(batch_size, dim=0)
# Model prediction
noise_scheduler.set_timesteps(steps)
for t in tqdm(noise_scheduler.timesteps):
latent_model_input = torch.cat([latents] * 2, dim=0)
latent_model_input = noise_scheduler.scale_model_input(latent_model_input, timestep=t)
noise_pred = unet(latent_model_input, t, encoder_hidden_states=t5_embeds_repeated).sample
noise_pred_uncond, noise_pred_text = noise_pred.chunk(2, dim=0)
noise_pred = noise_pred_uncond + guidance_scale * (noise_pred_text - noise_pred_uncond)
latents = noise_scheduler.step(noise_pred, t, latents).prev_sample
unet.to(offload_device)
# Decoding
vae.to(device)
latents = 1 / 0.18215 * latents
image = vae.decode(latents).sample
vae.to(offload_device)
image = (image / 2 + 0.5).clamp(0, 1)
image = image.permute(0, 2, 3, 1).cpu().float()
return (image,)
NODE_CLASS_MAPPINGS = {
"lavibridge_sampler": lavibridge_sampler,
"lavi_bridge_t5_encoder": lavi_bridge_t5_encoder,
"lavibridge_model_loader": lavibridge_model_loader
}
NODE_DISPLAY_NAME_MAPPINGS = {
"lavibridge_sampler": "LaVi-Bridge Sampler",
"lavi_bridge_t5_encoder": "LaVi-Bridge T5 Encoder",
"lavibridge_model_loader": "LaVi-Bridge Model Loader"
}
+3
View File
@@ -0,0 +1,3 @@
diffusers>=0.26.0
sentencepiece
peft>=0.8.2