diff --git a/.gitignore b/.gitignore new file mode 100644 index 0000000..e1cd816 --- /dev/null +++ b/.gitignore @@ -0,0 +1,9 @@ +pretrained_models/ +example_data/ +results/ +*.zip +.vscode/ +.hypothesis/ +*.pt +__pycache__ +*.pyc \ No newline at end of file diff --git a/LICENSE b/LICENSE new file mode 100644 index 0000000..7965606 --- /dev/null +++ b/LICENSE @@ -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 \ No newline at end of file diff --git a/README.md b/README.md new file mode 100644 index 0000000..79f4cce --- /dev/null +++ b/README.md @@ -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` \ No newline at end of file diff --git a/__init__.py b/__init__.py new file mode 100644 index 0000000..2e96bd6 --- /dev/null +++ b/__init__.py @@ -0,0 +1,3 @@ +from .nodes import NODE_CLASS_MAPPINGS, NODE_DISPLAY_NAME_MAPPINGS + +__all__ = ["NODE_CLASS_MAPPINGS", "NODE_DISPLAY_NAME_MAPPINGS"] \ No newline at end of file diff --git a/configs/v1-inference.yaml b/configs/v1-inference.yaml new file mode 100644 index 0000000..d4effe5 --- /dev/null +++ b/configs/v1-inference.yaml @@ -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 diff --git a/examples/lavi_t5_sd15_example_workflow.json b/examples/lavi_t5_sd15_example_workflow.json new file mode 100644 index 0000000..02efc97 --- /dev/null +++ b/examples/lavi_t5_sd15_example_workflow.json @@ -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 +} \ No newline at end of file diff --git a/modules/adapters.py b/modules/adapters.py new file mode 100644 index 0000000..f6e82f9 --- /dev/null +++ b/modules/adapters.py @@ -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) \ No newline at end of file diff --git a/modules/lora.py b/modules/lora.py new file mode 100644 index 0000000..2ead6e4 --- /dev/null +++ b/modules/lora.py @@ -0,0 +1,1106 @@ +import json +from itertools import groupby +from typing import Dict, List, Optional, Set, Tuple, Type, Union + +import torch +import torch.nn as nn + +try: + from safetensors.torch import safe_open + from safetensors.torch import save_file as safe_save + + safetensors_available = True +except ImportError: + from .safe_open import safe_open + + def safe_save( + tensors: Dict[str, torch.Tensor], + filename: str, + metadata: Optional[Dict[str, str]] = None, + ) -> None: + raise EnvironmentError( + "Saving safetensors requires the safetensors library. Please install with pip or similar." + ) + + safetensors_available = False + + +class LoraInjectedLinear(nn.Module): + def __init__( + self, in_features, out_features, bias=False, r=4, dropout_p=0.1, scale=1.0 + ): + super().__init__() + + if r > min(in_features, out_features): + raise ValueError( + f"LoRA rank {r} must be less or equal than {min(in_features, out_features)}" + ) + self.r = r + self.linear = nn.Linear(in_features, out_features, bias) + self.lora_down = nn.Linear(in_features, r, bias=False) + self.dropout = nn.Dropout(dropout_p) + self.lora_up = nn.Linear(r, out_features, bias=False) + self.scale = scale + self.selector = nn.Identity() + + nn.init.normal_(self.lora_down.weight, std=1 / r) + nn.init.zeros_(self.lora_up.weight) + + def forward(self, input): + return ( + self.linear(input) + + self.dropout(self.lora_up(self.selector(self.lora_down(input)))) + * self.scale + ) + + def realize_as_lora(self): + return self.lora_up.weight.data * self.scale, self.lora_down.weight.data + + def set_selector_from_diag(self, diag: torch.Tensor): + # diag is a 1D tensor of size (r,) + assert diag.shape == (self.r,) + self.selector = nn.Linear(self.r, self.r, bias=False) + self.selector.weight.data = torch.diag(diag) + self.selector.weight.data = self.selector.weight.data.to( + self.lora_up.weight.device + ).to(self.lora_up.weight.dtype) + + +class LoraInjectedConv2d(nn.Module): + def __init__( + self, + in_channels: int, + out_channels: int, + kernel_size, + stride=1, + padding=0, + dilation=1, + groups: int = 1, + bias: bool = True, + r: int = 4, + dropout_p: float = 0.1, + scale: float = 1.0, + ): + super().__init__() + if r > min(in_channels, out_channels): + raise ValueError( + f"LoRA rank {r} must be less or equal than {min(in_channels, out_channels)}" + ) + self.r = r + self.conv = nn.Conv2d( + in_channels=in_channels, + out_channels=out_channels, + kernel_size=kernel_size, + stride=stride, + padding=padding, + dilation=dilation, + groups=groups, + bias=bias, + ) + + self.lora_down = nn.Conv2d( + in_channels=in_channels, + out_channels=r, + kernel_size=kernel_size, + stride=stride, + padding=padding, + dilation=dilation, + groups=groups, + bias=False, + ) + self.dropout = nn.Dropout(dropout_p) + self.lora_up = nn.Conv2d( + in_channels=r, + out_channels=out_channels, + kernel_size=1, + stride=1, + padding=0, + bias=False, + ) + self.selector = nn.Identity() + self.scale = scale + + nn.init.normal_(self.lora_down.weight, std=1 / r) + nn.init.zeros_(self.lora_up.weight) + + def forward(self, input): + return ( + self.conv(input) + + self.dropout(self.lora_up(self.selector(self.lora_down(input)))) + * self.scale + ) + + def realize_as_lora(self): + return self.lora_up.weight.data * self.scale, self.lora_down.weight.data + + def set_selector_from_diag(self, diag: torch.Tensor): + # diag is a 1D tensor of size (r,) + assert diag.shape == (self.r,) + self.selector = nn.Conv2d( + in_channels=self.r, + out_channels=self.r, + kernel_size=1, + stride=1, + padding=0, + bias=False, + ) + self.selector.weight.data = torch.diag(diag) + + # same device + dtype as lora_up + self.selector.weight.data = self.selector.weight.data.to( + self.lora_up.weight.device + ).to(self.lora_up.weight.dtype) + + +UNET_DEFAULT_TARGET_REPLACE = {"CrossAttention", "Attention", "GEGLU"} + +UNET_EXTENDED_TARGET_REPLACE = {"ResnetBlock2D", "CrossAttention", "Attention", "GEGLU"} + +TEXT_ENCODER_DEFAULT_TARGET_REPLACE = {"CLIPAttention"} + +TEXT_ENCODER_EXTENDED_TARGET_REPLACE = {"CLIPAttention"} + +DEFAULT_TARGET_REPLACE = UNET_DEFAULT_TARGET_REPLACE + +EMBED_FLAG = "" + + +def _find_children( + model, + search_class: List[Type[nn.Module]] = [nn.Linear], +): + """ + Find all modules of a certain class (or union of classes). + + Returns all matching modules, along with the parent of those moduless and the + names they are referenced by. + """ + # For each target find every linear_class module that isn't a child of a LoraInjectedLinear + for parent in model.modules(): + for name, module in parent.named_children(): + if any([isinstance(module, _class) for _class in search_class]): + yield parent, name, module + + +def _find_modules_v2( + model, + ancestor_class: Optional[Set[str]] = None, + search_class: List[Type[nn.Module]] = [nn.Linear], + exclude_children_of: Optional[List[Type[nn.Module]]] = [ + LoraInjectedLinear, + LoraInjectedConv2d, + ], +): + """ + Find all modules of a certain class (or union of classes) that are direct or + indirect descendants of other modules of a certain class (or union of classes). + + Returns all matching modules, along with the parent of those moduless and the + names they are referenced by. + """ + + # Get the targets we should replace all linears under + if ancestor_class is not None: + ancestors = ( + module + for module in model.modules() + if module.__class__.__name__ in ancestor_class + ) + else: + # this, incase you want to naively iterate over all modules. + ancestors = [module for module in model.modules()] + + # For each target find every linear_class module that isn't a child of a LoraInjectedLinear + for ancestor in ancestors: + for fullname, module in ancestor.named_modules(): + if any([isinstance(module, _class) for _class in search_class]): + # Find the direct parent if this is a descendant, not a child, of target + *path, name = fullname.split(".") + parent = ancestor + while path: + parent = parent.get_submodule(path.pop(0)) + # Skip this linear if it's a child of a LoraInjectedLinear + if exclude_children_of and any( + [isinstance(parent, _class) for _class in exclude_children_of] + ): + continue + # Otherwise, yield it + yield parent, name, module + + +def _find_modules_old( + model, + ancestor_class: Set[str] = DEFAULT_TARGET_REPLACE, + search_class: List[Type[nn.Module]] = [nn.Linear], + exclude_children_of: Optional[List[Type[nn.Module]]] = [LoraInjectedLinear], +): + ret = [] + for _module in model.modules(): + if _module.__class__.__name__ in ancestor_class: + + for name, _child_module in _module.named_modules(): + if _child_module.__class__ in search_class: + ret.append((_module, name, _child_module)) + print(ret) + return ret + + +_find_modules = _find_modules_v2 + + +def inject_trainable_lora( + model: nn.Module, + target_replace_module: Set[str] = DEFAULT_TARGET_REPLACE, + r: int = 4, + loras=None, # path to lora .pt + verbose: bool = False, + dropout_p: float = 0.0, + scale: float = 1.0, +): + """ + inject lora into model, and returns lora parameter groups. + """ + + require_grad_params = [] + names = [] + + if loras != None: + loras = torch.load(loras) + + for _module, name, _child_module in _find_modules( + model, target_replace_module, search_class=[nn.Linear] + ): + weight = _child_module.weight + bias = _child_module.bias + if verbose: + print("LoRA Injection : injecting lora into ", name) + print("LoRA Injection : weight shape", weight.shape) + _tmp = LoraInjectedLinear( + _child_module.in_features, + _child_module.out_features, + _child_module.bias is not None, + r=r, + dropout_p=dropout_p, + scale=scale, + ) + _tmp.linear.weight = weight + if bias is not None: + _tmp.linear.bias = bias + + # switch the module + _tmp.to(_child_module.weight.device).to(_child_module.weight.dtype) + _module._modules[name] = _tmp + + require_grad_params.append(_module._modules[name].lora_up.parameters()) + require_grad_params.append(_module._modules[name].lora_down.parameters()) + + if loras != None: + _module._modules[name].lora_up.weight = loras.pop(0) + _module._modules[name].lora_down.weight = loras.pop(0) + + _module._modules[name].lora_up.weight.requires_grad = True + _module._modules[name].lora_down.weight.requires_grad = True + names.append(name) + + return require_grad_params, names + + +def inject_trainable_lora_extended( + model: nn.Module, + target_replace_module: Set[str] = UNET_EXTENDED_TARGET_REPLACE, + r: int = 4, + loras=None, # path to lora .pt +): + """ + inject lora into model, and returns lora parameter groups. + """ + + require_grad_params = [] + names = [] + + if loras != None: + loras = torch.load(loras) + + for _module, name, _child_module in _find_modules( + model, target_replace_module, search_class=[nn.Linear, nn.Conv2d] + ): + if _child_module.__class__ == nn.Linear: + weight = _child_module.weight + bias = _child_module.bias + _tmp = LoraInjectedLinear( + _child_module.in_features, + _child_module.out_features, + _child_module.bias is not None, + r=r, + ) + _tmp.linear.weight = weight + if bias is not None: + _tmp.linear.bias = bias + elif _child_module.__class__ == nn.Conv2d: + weight = _child_module.weight + bias = _child_module.bias + _tmp = LoraInjectedConv2d( + _child_module.in_channels, + _child_module.out_channels, + _child_module.kernel_size, + _child_module.stride, + _child_module.padding, + _child_module.dilation, + _child_module.groups, + _child_module.bias is not None, + r=r, + ) + + _tmp.conv.weight = weight + if bias is not None: + _tmp.conv.bias = bias + + # switch the module + _tmp.to(_child_module.weight.device).to(_child_module.weight.dtype) + if bias is not None: + _tmp.to(_child_module.bias.device).to(_child_module.bias.dtype) + + _module._modules[name] = _tmp + + require_grad_params.append(_module._modules[name].lora_up.parameters()) + require_grad_params.append(_module._modules[name].lora_down.parameters()) + + if loras != None: + _module._modules[name].lora_up.weight = loras.pop(0) + _module._modules[name].lora_down.weight = loras.pop(0) + + _module._modules[name].lora_up.weight.requires_grad = True + _module._modules[name].lora_down.weight.requires_grad = True + names.append(name) + + return require_grad_params, names + + +def extract_lora_ups_down(model, target_replace_module=DEFAULT_TARGET_REPLACE): + + loras = [] + + for _m, _n, _child_module in _find_modules( + model, + target_replace_module, + search_class=[LoraInjectedLinear, LoraInjectedConv2d], + ): + loras.append((_child_module.lora_up, _child_module.lora_down)) + + if len(loras) == 0: + raise ValueError("No lora injected.") + + return loras + + +def extract_lora_as_tensor( + model, target_replace_module=DEFAULT_TARGET_REPLACE, as_fp16=True +): + + loras = [] + + for _m, _n, _child_module in _find_modules( + model, + target_replace_module, + search_class=[LoraInjectedLinear, LoraInjectedConv2d], + ): + up, down = _child_module.realize_as_lora() + if as_fp16: + up = up.to(torch.float16) + down = down.to(torch.float16) + + loras.append((up, down)) + + if len(loras) == 0: + raise ValueError("No lora injected.") + + return loras + + +def save_lora_weight( + model, + path="./lora.pt", + target_replace_module=DEFAULT_TARGET_REPLACE, +): + weights = [] + for _up, _down in extract_lora_ups_down( + model, target_replace_module=target_replace_module + ): + weights.append(_up.weight.to("cpu").to(torch.float16)) + weights.append(_down.weight.to("cpu").to(torch.float16)) + + torch.save(weights, path) + + +def save_lora_as_json(model, path="./lora.json"): + weights = [] + for _up, _down in extract_lora_ups_down(model): + weights.append(_up.weight.detach().cpu().numpy().tolist()) + weights.append(_down.weight.detach().cpu().numpy().tolist()) + + import json + + with open(path, "w") as f: + json.dump(weights, f) + + +def save_safeloras_with_embeds( + modelmap: Dict[str, Tuple[nn.Module, Set[str]]] = {}, + embeds: Dict[str, torch.Tensor] = {}, + outpath="./lora.safetensors", +): + """ + Saves the Lora from multiple modules in a single safetensor file. + + modelmap is a dictionary of { + "module name": (module, target_replace_module) + } + """ + weights = {} + metadata = {} + + for name, (model, target_replace_module) in modelmap.items(): + metadata[name] = json.dumps(list(target_replace_module)) + + for i, (_up, _down) in enumerate( + extract_lora_as_tensor(model, target_replace_module) + ): + rank = _down.shape[0] + + metadata[f"{name}:{i}:rank"] = str(rank) + weights[f"{name}:{i}:up"] = _up + weights[f"{name}:{i}:down"] = _down + + for token, tensor in embeds.items(): + metadata[token] = EMBED_FLAG + weights[token] = tensor + + print(f"Saving weights to {outpath}") + safe_save(weights, outpath, metadata) + + +def save_safeloras( + modelmap: Dict[str, Tuple[nn.Module, Set[str]]] = {}, + outpath="./lora.safetensors", +): + return save_safeloras_with_embeds(modelmap=modelmap, outpath=outpath) + + +def convert_loras_to_safeloras_with_embeds( + modelmap: Dict[str, Tuple[str, Set[str], int]] = {}, + embeds: Dict[str, torch.Tensor] = {}, + outpath="./lora.safetensors", +): + """ + Converts the Lora from multiple pytorch .pt files into a single safetensor file. + + modelmap is a dictionary of { + "module name": (pytorch_model_path, target_replace_module, rank) + } + """ + + weights = {} + metadata = {} + + for name, (path, target_replace_module, r) in modelmap.items(): + metadata[name] = json.dumps(list(target_replace_module)) + + lora = torch.load(path) + for i, weight in enumerate(lora): + is_up = i % 2 == 0 + i = i // 2 + + if is_up: + metadata[f"{name}:{i}:rank"] = str(r) + weights[f"{name}:{i}:up"] = weight + else: + weights[f"{name}:{i}:down"] = weight + + for token, tensor in embeds.items(): + metadata[token] = EMBED_FLAG + weights[token] = tensor + + print(f"Saving weights to {outpath}") + safe_save(weights, outpath, metadata) + + +def convert_loras_to_safeloras( + modelmap: Dict[str, Tuple[str, Set[str], int]] = {}, + outpath="./lora.safetensors", +): + convert_loras_to_safeloras_with_embeds(modelmap=modelmap, outpath=outpath) + + +def parse_safeloras( + safeloras, +) -> Dict[str, Tuple[List[nn.parameter.Parameter], List[int], List[str]]]: + """ + Converts a loaded safetensor file that contains a set of module Loras + into Parameters and other information + + Output is a dictionary of { + "module name": ( + [list of weights], + [list of ranks], + target_replacement_modules + ) + } + """ + loras = {} + metadata = safeloras.metadata() + + get_name = lambda k: k.split(":")[0] + + keys = list(safeloras.keys()) + keys.sort(key=get_name) + + for name, module_keys in groupby(keys, get_name): + info = metadata.get(name) + + if not info: + raise ValueError( + f"Tensor {name} has no metadata - is this a Lora safetensor?" + ) + + # Skip Textual Inversion embeds + if info == EMBED_FLAG: + continue + + # Handle Loras + # Extract the targets + target = json.loads(info) + + # Build the result lists - Python needs us to preallocate lists to insert into them + module_keys = list(module_keys) + ranks = [4] * (len(module_keys) // 2) + weights = [None] * len(module_keys) + + for key in module_keys: + # Split the model name and index out of the key + _, idx, direction = key.split(":") + idx = int(idx) + + # Add the rank + ranks[idx] = int(metadata[f"{name}:{idx}:rank"]) + + # Insert the weight into the list + idx = idx * 2 + (1 if direction == "down" else 0) + weights[idx] = nn.parameter.Parameter(safeloras.get_tensor(key)) + + loras[name] = (weights, ranks, target) + + return loras + + +def parse_safeloras_embeds( + safeloras, +) -> Dict[str, torch.Tensor]: + """ + Converts a loaded safetensor file that contains Textual Inversion embeds into + a dictionary of embed_token: Tensor + """ + embeds = {} + metadata = safeloras.metadata() + + for key in safeloras.keys(): + # Only handle Textual Inversion embeds + meta = metadata.get(key) + if not meta or meta != EMBED_FLAG: + continue + + embeds[key] = safeloras.get_tensor(key) + + return embeds + + +def load_safeloras(path, device="cpu"): + safeloras = safe_open(path, framework="pt", device=device) + return parse_safeloras(safeloras) + + +def load_safeloras_embeds(path, device="cpu"): + safeloras = safe_open(path, framework="pt", device=device) + return parse_safeloras_embeds(safeloras) + + +def load_safeloras_both(path, device="cpu"): + safeloras = safe_open(path, framework="pt", device=device) + return parse_safeloras(safeloras), parse_safeloras_embeds(safeloras) + + +def collapse_lora(model, alpha=1.0): + + for _module, name, _child_module in _find_modules( + model, + UNET_EXTENDED_TARGET_REPLACE | TEXT_ENCODER_EXTENDED_TARGET_REPLACE, + search_class=[LoraInjectedLinear, LoraInjectedConv2d], + ): + + if isinstance(_child_module, LoraInjectedLinear): + print("Collapsing Lin Lora in", name) + + _child_module.linear.weight = nn.Parameter( + _child_module.linear.weight.data + + alpha + * ( + _child_module.lora_up.weight.data + @ _child_module.lora_down.weight.data + ) + .type(_child_module.linear.weight.dtype) + .to(_child_module.linear.weight.device) + ) + + else: + print("Collapsing Conv Lora in", name) + _child_module.conv.weight = nn.Parameter( + _child_module.conv.weight.data + + alpha + * ( + _child_module.lora_up.weight.data.flatten(start_dim=1) + @ _child_module.lora_down.weight.data.flatten(start_dim=1) + ) + .reshape(_child_module.conv.weight.data.shape) + .type(_child_module.conv.weight.dtype) + .to(_child_module.conv.weight.device) + ) + + +def monkeypatch_or_replace_lora( + model, + loras, + target_replace_module=DEFAULT_TARGET_REPLACE, + r: Union[int, List[int]] = 4, +): + for _module, name, _child_module in _find_modules( + model, target_replace_module, search_class=[nn.Linear, LoraInjectedLinear] + ): + _source = ( + _child_module.linear + if isinstance(_child_module, LoraInjectedLinear) + else _child_module + ) + + weight = _source.weight + bias = _source.bias + _tmp = LoraInjectedLinear( + _source.in_features, + _source.out_features, + _source.bias is not None, + r=r.pop(0) if isinstance(r, list) else r, + ) + _tmp.linear.weight = weight + + if bias is not None: + _tmp.linear.bias = bias + + # switch the module + _module._modules[name] = _tmp + + up_weight = loras.pop(0) + down_weight = loras.pop(0) + + _module._modules[name].lora_up.weight = nn.Parameter( + up_weight.type(weight.dtype) + ) + _module._modules[name].lora_down.weight = nn.Parameter( + down_weight.type(weight.dtype) + ) + + _module._modules[name].to(weight.device) + + +def monkeypatch_or_replace_lora_extended( + model, + loras, + target_replace_module=DEFAULT_TARGET_REPLACE, + r: Union[int, List[int]] = 4, +): + for _module, name, _child_module in _find_modules( + model, + target_replace_module, + search_class=[nn.Linear, LoraInjectedLinear, nn.Conv2d, LoraInjectedConv2d], + ): + + if (_child_module.__class__ == nn.Linear) or ( + _child_module.__class__ == LoraInjectedLinear + ): + if len(loras[0].shape) != 2: + continue + + _source = ( + _child_module.linear + if isinstance(_child_module, LoraInjectedLinear) + else _child_module + ) + + weight = _source.weight + bias = _source.bias + _tmp = LoraInjectedLinear( + _source.in_features, + _source.out_features, + _source.bias is not None, + r=r.pop(0) if isinstance(r, list) else r, + ) + _tmp.linear.weight = weight + + if bias is not None: + _tmp.linear.bias = bias + + elif (_child_module.__class__ == nn.Conv2d) or ( + _child_module.__class__ == LoraInjectedConv2d + ): + if len(loras[0].shape) != 4: + continue + _source = ( + _child_module.conv + if isinstance(_child_module, LoraInjectedConv2d) + else _child_module + ) + + weight = _source.weight + bias = _source.bias + _tmp = LoraInjectedConv2d( + _source.in_channels, + _source.out_channels, + _source.kernel_size, + _source.stride, + _source.padding, + _source.dilation, + _source.groups, + _source.bias is not None, + r=r.pop(0) if isinstance(r, list) else r, + ) + + _tmp.conv.weight = weight + + if bias is not None: + _tmp.conv.bias = bias + + # switch the module + _module._modules[name] = _tmp + + up_weight = loras.pop(0) + down_weight = loras.pop(0) + + _module._modules[name].lora_up.weight = nn.Parameter( + up_weight.type(weight.dtype) + ) + _module._modules[name].lora_down.weight = nn.Parameter( + down_weight.type(weight.dtype) + ) + + _module._modules[name].to(weight.device) + + +def monkeypatch_or_replace_safeloras(models, safeloras): + loras = parse_safeloras(safeloras) + + for name, (lora, ranks, target) in loras.items(): + model = getattr(models, name, None) + + if not model: + print(f"No model provided for {name}, contained in Lora") + continue + + monkeypatch_or_replace_lora_extended(model, lora, target, ranks) + + +def monkeypatch_remove_lora(model): + for _module, name, _child_module in _find_modules( + model, search_class=[LoraInjectedLinear, LoraInjectedConv2d] + ): + if isinstance(_child_module, LoraInjectedLinear): + _source = _child_module.linear + weight, bias = _source.weight, _source.bias + + _tmp = nn.Linear( + _source.in_features, _source.out_features, bias is not None + ) + + _tmp.weight = weight + if bias is not None: + _tmp.bias = bias + + else: + _source = _child_module.conv + weight, bias = _source.weight, _source.bias + + _tmp = nn.Conv2d( + in_channels=_source.in_channels, + out_channels=_source.out_channels, + kernel_size=_source.kernel_size, + stride=_source.stride, + padding=_source.padding, + dilation=_source.dilation, + groups=_source.groups, + bias=bias is not None, + ) + + _tmp.weight = weight + if bias is not None: + _tmp.bias = bias + + _module._modules[name] = _tmp + + +def monkeypatch_add_lora( + model, + loras, + target_replace_module=DEFAULT_TARGET_REPLACE, + alpha: float = 1.0, + beta: float = 1.0, +): + for _module, name, _child_module in _find_modules( + model, target_replace_module, search_class=[LoraInjectedLinear] + ): + weight = _child_module.linear.weight + + up_weight = loras.pop(0) + down_weight = loras.pop(0) + + _module._modules[name].lora_up.weight = nn.Parameter( + up_weight.type(weight.dtype).to(weight.device) * alpha + + _module._modules[name].lora_up.weight.to(weight.device) * beta + ) + _module._modules[name].lora_down.weight = nn.Parameter( + down_weight.type(weight.dtype).to(weight.device) * alpha + + _module._modules[name].lora_down.weight.to(weight.device) * beta + ) + + _module._modules[name].to(weight.device) + + +def tune_lora_scale(model, alpha: float = 1.0): + for _module in model.modules(): + if _module.__class__.__name__ in ["LoraInjectedLinear", "LoraInjectedConv2d"]: + _module.scale = alpha + + +def set_lora_diag(model, diag: torch.Tensor): + for _module in model.modules(): + if _module.__class__.__name__ in ["LoraInjectedLinear", "LoraInjectedConv2d"]: + _module.set_selector_from_diag(diag) + + +def _text_lora_path(path: str) -> str: + assert path.endswith(".pt"), "Only .pt files are supported" + return ".".join(path.split(".")[:-1] + ["text_encoder", "pt"]) + + +def _ti_lora_path(path: str) -> str: + assert path.endswith(".pt"), "Only .pt files are supported" + return ".".join(path.split(".")[:-1] + ["ti", "pt"]) + + +def apply_learned_embed_in_clip( + learned_embeds, + text_encoder, + tokenizer, + token: Optional[Union[str, List[str]]] = None, + idempotent=False, +): + if isinstance(token, str): + trained_tokens = [token] + elif isinstance(token, list): + assert len(learned_embeds.keys()) == len( + token + ), "The number of tokens and the number of embeds should be the same" + trained_tokens = token + else: + trained_tokens = list(learned_embeds.keys()) + + for token in trained_tokens: + print(token) + embeds = learned_embeds[token] + + # cast to dtype of text_encoder + dtype = text_encoder.get_input_embeddings().weight.dtype + num_added_tokens = tokenizer.add_tokens(token) + + i = 1 + if not idempotent: + while num_added_tokens == 0: + print(f"The tokenizer already contains the token {token}.") + token = f"{token[:-1]}-{i}>" + print(f"Attempting to add the token {token}.") + num_added_tokens = tokenizer.add_tokens(token) + i += 1 + elif num_added_tokens == 0 and idempotent: + print(f"The tokenizer already contains the token {token}.") + print(f"Replacing {token} embedding.") + + # resize the token embeddings + text_encoder.resize_token_embeddings(len(tokenizer)) + + # get the id for the token and assign the embeds + token_id = tokenizer.convert_tokens_to_ids(token) + text_encoder.get_input_embeddings().weight.data[token_id] = embeds + return token + + +def load_learned_embed_in_clip( + learned_embeds_path, + text_encoder, + tokenizer, + token: Optional[Union[str, List[str]]] = None, + idempotent=False, +): + learned_embeds = torch.load(learned_embeds_path) + apply_learned_embed_in_clip( + learned_embeds, text_encoder, tokenizer, token, idempotent + ) + + +def patch_pipe( + pipe, + maybe_unet_path, + token: Optional[str] = None, + r: int = 4, + patch_unet=True, + patch_text=True, + patch_ti=True, + idempotent_token=True, + unet_target_replace_module=DEFAULT_TARGET_REPLACE, + text_target_replace_module=TEXT_ENCODER_DEFAULT_TARGET_REPLACE, +): + if maybe_unet_path.endswith(".pt"): + # torch format + + if maybe_unet_path.endswith(".ti.pt"): + unet_path = maybe_unet_path[:-6] + ".pt" + elif maybe_unet_path.endswith(".text_encoder.pt"): + unet_path = maybe_unet_path[:-16] + ".pt" + else: + unet_path = maybe_unet_path + + ti_path = _ti_lora_path(unet_path) + text_path = _text_lora_path(unet_path) + + if patch_unet: + print("LoRA : Patching Unet") + monkeypatch_or_replace_lora( + pipe.unet, + torch.load(unet_path), + r=r, + target_replace_module=unet_target_replace_module, + ) + + if patch_text: + print("LoRA : Patching text encoder") + monkeypatch_or_replace_lora( + pipe.text_encoder, + torch.load(text_path), + target_replace_module=text_target_replace_module, + r=r, + ) + if patch_ti: + print("LoRA : Patching token input") + token = load_learned_embed_in_clip( + ti_path, + pipe.text_encoder, + pipe.tokenizer, + token=token, + idempotent=idempotent_token, + ) + + elif maybe_unet_path.endswith(".safetensors"): + safeloras = safe_open(maybe_unet_path, framework="pt", device="cpu") + monkeypatch_or_replace_safeloras(pipe, safeloras) + tok_dict = parse_safeloras_embeds(safeloras) + if patch_ti: + apply_learned_embed_in_clip( + tok_dict, + pipe.text_encoder, + pipe.tokenizer, + token=token, + idempotent=idempotent_token, + ) + return tok_dict + + +@torch.no_grad() +def inspect_lora(model): + moved = {} + + for name, _module in model.named_modules(): + if _module.__class__.__name__ in ["LoraInjectedLinear", "LoraInjectedConv2d"]: + ups = _module.lora_up.weight.data.clone() + downs = _module.lora_down.weight.data.clone() + + wght: torch.Tensor = ups.flatten(1) @ downs.flatten(1) + + dist = wght.flatten().abs().mean().item() + if name in moved: + moved[name].append(dist) + else: + moved[name] = [dist] + + return moved + + +def save_all( + unet, + text_encoder, + save_path, + placeholder_token_ids=None, + placeholder_tokens=None, + save_lora=True, + save_ti=True, + target_replace_module_text=TEXT_ENCODER_DEFAULT_TARGET_REPLACE, + target_replace_module_unet=DEFAULT_TARGET_REPLACE, + safe_form=True, +): + if not safe_form: + # save ti + if save_ti: + ti_path = _ti_lora_path(save_path) + learned_embeds_dict = {} + for tok, tok_id in zip(placeholder_tokens, placeholder_token_ids): + learned_embeds = text_encoder.get_input_embeddings().weight[tok_id] + print( + f"Current Learned Embeddings for {tok}:, id {tok_id} ", + learned_embeds[:4], + ) + learned_embeds_dict[tok] = learned_embeds.detach().cpu() + + torch.save(learned_embeds_dict, ti_path) + print("Ti saved to ", ti_path) + + # save text encoder + if save_lora: + + save_lora_weight( + unet, save_path, target_replace_module=target_replace_module_unet + ) + print("Unet saved to ", save_path) + + save_lora_weight( + text_encoder, + _text_lora_path(save_path), + target_replace_module=target_replace_module_text, + ) + print("Text Encoder saved to ", _text_lora_path(save_path)) + + else: + assert save_path.endswith( + ".safetensors" + ), f"Save path : {save_path} should end with .safetensors" + + loras = {} + embeds = {} + + if save_lora: + + loras["unet"] = (unet, target_replace_module_unet) + loras["text_encoder"] = (text_encoder, target_replace_module_text) + + if save_ti: + for tok, tok_id in zip(placeholder_tokens, placeholder_token_ids): + learned_embeds = text_encoder.get_input_embeddings().weight[tok_id] + print( + f"Current Learned Embeddings for {tok}:, id {tok_id} ", + learned_embeds[:4], + ) + embeds[tok] = learned_embeds.detach().cpu() + + save_safeloras_with_embeds(loras, embeds, save_path) diff --git a/nodes.py b/nodes.py new file mode 100644 index 0000000..97cd120 --- /dev/null +++ b/nodes.py @@ -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" +} diff --git a/requirements.txt b/requirements.txt new file mode 100644 index 0000000..1fb0977 --- /dev/null +++ b/requirements.txt @@ -0,0 +1,3 @@ +diffusers>=0.26.0 +sentencepiece +peft>=0.8.2 \ No newline at end of file