Initial commit
This commit is contained in:
@@ -0,0 +1,9 @@
|
||||
pretrained_models/
|
||||
example_data/
|
||||
results/
|
||||
*.zip
|
||||
.vscode/
|
||||
.hypothesis/
|
||||
*.pt
|
||||
__pycache__
|
||||
*.pyc
|
||||
@@ -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
|
||||
@@ -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`
|
||||
@@ -0,0 +1,3 @@
|
||||
from .nodes import NODE_CLASS_MAPPINGS, NODE_DISPLAY_NAME_MAPPINGS
|
||||
|
||||
__all__ = ["NODE_CLASS_MAPPINGS", "NODE_DISPLAY_NAME_MAPPINGS"]
|
||||
@@ -0,0 +1,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
|
||||
@@ -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
|
||||
}
|
||||
@@ -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
File diff suppressed because it is too large
Load Diff
@@ -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"
|
||||
}
|
||||
@@ -0,0 +1,3 @@
|
||||
diffusers>=0.26.0
|
||||
sentencepiece
|
||||
peft>=0.8.2
|
||||
Reference in New Issue
Block a user