diff --git a/py/config.py b/py/config.py index e08773e..0de9be4 100644 --- a/py/config.py +++ b/py/config.py @@ -61,30 +61,67 @@ FOOOCUS_INPAINT_PATCH = { LAYER_DIFFUSION_VAE = { "encode": { - "model_url": "https://huggingface.co/LayerDiffusion/layerdiffusion-v1/resolve/main/vae_transparent_encoder.safetensors" + "sdxl": { + "model_url": "https://huggingface.co/LayerDiffusion/layerdiffusion-v1/resolve/main/vae_transparent_encoder.safetensors" + } }, "decode": { - "model_url": "https://huggingface.co/LayerDiffusion/layerdiffusion-v1/resolve/main/vae_transparent_decoder.safetensors" + "sd15": { + "model_url": "https://huggingface.co/LayerDiffusion/layerdiffusion-v1/resolve/main/layer_sd15_vae_transparent_decoder.safetensors" + }, + "sdxl": { + "model_url": "https://huggingface.co/LayerDiffusion/layerdiffusion-v1/resolve/main/vae_transparent_decoder.safetensors" + } } } LAYER_DIFFUSION = { "Attention Injection": { - "model_url": "https://huggingface.co/LayerDiffusion/layerdiffusion-v1/resolve/main/layer_xl_transparent_attn.safetensors" + "sd15": { + "model_url": "https://huggingface.co/LayerDiffusion/layerdiffusion-v1/resolve/main/layer_sd15_transparent_attn.safetensors" + }, + "sdxl": { + "model_url": "https://huggingface.co/LayerDiffusion/layerdiffusion-v1/resolve/main/layer_xl_transparent_attn.safetensors" + }, }, "Conv Injection": { - "model_url": "https://huggingface.co/LayerDiffusion/layerdiffusion-v1/resolve/main/layer_xl_transparent_conv.safetensors" + "sdxl": { + "model_url": "https://huggingface.co/LayerDiffusion/layerdiffusion-v1/resolve/main/layer_xl_transparent_conv.safetensors" + }, + "sd15": { + "model_url": None + } }, "Foreground": { - "model_url": "https://huggingface.co/LayerDiffusion/layerdiffusion-v1/resolve/main/layer_xl_fg2ble.safetensors" + "sd15": { + "model_url": "https://huggingface.co/LayerDiffusion/layerdiffusion-v1/resolve/main/layer_sd15_joint.safetensors" + }, + "sdxl": { + "model_url": "https://huggingface.co/LayerDiffusion/layerdiffusion-v1/resolve/main/layer_xl_fg2ble.safetensors" + } }, "Foreground to Background": { - "model_url": "https://huggingface.co/LayerDiffusion/layerdiffusion-v1/resolve/main/layer_xl_fgble2bg.safetensors" + "sd15": { + "model_url": "https://huggingface.co/LayerDiffusion/layerdiffusion-v1/resolve/main/layer_sd15_fg2bg.safetensors" + }, + "sdxl": { + "model_url": "https://huggingface.co/LayerDiffusion/layerdiffusion-v1/resolve/main/layer_xl_fgble2bg.safetensors" + } }, "Background": { - "model_url": "https://huggingface.co/LayerDiffusion/layerdiffusion-v1/resolve/main/layer_xl_bg2ble.safetensors" + "sd15": { + "model_url": "https://huggingface.co/LayerDiffusion/layerdiffusion-v1/resolve/main/layer_sd15_joint.safetensors" + }, + "sdxl": { + "model_url": "https://huggingface.co/LayerDiffusion/layerdiffusion-v1/resolve/main/layer_xl_bg2ble.safetensors" + } }, "Background to Foreground": { - "model_url": "https://huggingface.co/LayerDiffusion/layerdiffusion-v1/resolve/main/layer_xl_bgble2fg.safetensors" + "sd15": { + "model_url": "https://huggingface.co/LayerDiffusion/layerdiffusion-v1/resolve/main/layer_sd15_bg2fg.safetensors" + }, + "sdxl": { + "model_url": "https://huggingface.co/LayerDiffusion/layerdiffusion-v1/resolve/main/layer_xl_bgble2fg.safetensor" + } }, } \ No newline at end of file diff --git a/py/easyNodes.py b/py/easyNodes.py index 39e9847..dcbfa57 100644 --- a/py/easyNodes.py +++ b/py/easyNodes.py @@ -15,7 +15,7 @@ from .config import MAX_SEED_NUM, BASE_RESOLUTIONS, RESOURCES_DIR, INPAINT_DIR, from .log import log_node_info, log_node_error, log_node_warn from .wildcards import process_with_loras, get_wildcard_list, process from .adv_encode import advanced_encode -from .layer_diffusion import LayerDiffuse, LayerMethod, calculate_weight_adjust_channel +from .layer_diffuse.func import LayerDiffuse, LayerMethod from .libs.utils import find_wildcards_seed, is_linked_styles_selector, easySave, get_local_filepath, add_folder_path_and_extensions from .libs.loader import easyLoader @@ -1687,7 +1687,7 @@ class instantIDApplyAdvanced(instantID): RETURN_NAMES = ("pipe", "model", "positive", "negative") OUTPUT_NODE = True - FUNCTION = "apply" + FUNCTION = "apply_advanced" CATEGORY = "EasyUse/__for_testing" def apply_advanced(self, pipe, image, instantid_file, insightface, control_net_name, cn_strength, cn_soft_weights, weight, start_at, end_at, noise, image_kps=None, mask=None, control_net=None, positive=None, negative=None, prompt=None, extra_pnginfo=None, my_unique_id=None): @@ -2524,10 +2524,6 @@ class samplerFull(LayerDiffuse): method = self.get_layer_diffusion_method(pipe['loader_settings']['layer_diffusion_method'], samp_blend_samples is not None) weight = pipe['loader_settings']['layer_diffusion_weight'] if 'layer_diffusion_weight' in pipe['loader_settings'] else 1.0 - try: - ModelPatcher.calculate_weight = calculate_weight_adjust_channel(ModelPatcher.calculate_weight) - except: - pass samp_model, samp_positive, samp_negative = self.apply_layer_diffusion(samp_model, method, weight, samp_samples, samp_blend_samples, samp_positive, samp_negative) def downscale_model_unet(samp_model): @@ -2594,7 +2590,7 @@ class samplerFull(LayerDiffuse): samp_images = samp_vae.decode(latent).cpu() # LayerDiffusion Decode - new_images, samp_images, alpha = self.layer_diffusion_decode(layer_diffusion_method, latent, blend_samples, samp_images) + new_images, samp_images, alpha = self.layer_diffusion_decode(layer_diffusion_method, latent, blend_samples, samp_images, samp_model) # 推理总耗时(包含解码) end_decode_time = int(time.time() * 1000) @@ -2715,7 +2711,7 @@ class samplerFull(LayerDiffuse): output_images = torch.stack([tensor.squeeze() for tensor in image_list]) new_images, samp_images, alpha = self.layer_diffusion_decode(layer_diffusion_method, latents_plot, blend_samples, - output_images) + output_images, samp_model) results = easySave(images, save_prefix, image_output, prompt, extra_pnginfo) sampler.update_value_by_id("results", my_unique_id, results) diff --git a/py/layer_diffuse/__init__.py b/py/layer_diffuse/__init__.py new file mode 100644 index 0000000..e69de29 diff --git a/py/layer_diffuse/attension_sharing.py b/py/layer_diffuse/attension_sharing.py new file mode 100644 index 0000000..2ed12da --- /dev/null +++ b/py/layer_diffuse/attension_sharing.py @@ -0,0 +1,360 @@ +# Currently only sd15 + +import functools +import torch +import einops + +from comfy import model_management, utils +from comfy.ldm.modules.attention import optimized_attention + + +module_mapping_sd15 = { + 0: "input_blocks.1.1.transformer_blocks.0.attn1", + 1: "input_blocks.1.1.transformer_blocks.0.attn2", + 2: "input_blocks.2.1.transformer_blocks.0.attn1", + 3: "input_blocks.2.1.transformer_blocks.0.attn2", + 4: "input_blocks.4.1.transformer_blocks.0.attn1", + 5: "input_blocks.4.1.transformer_blocks.0.attn2", + 6: "input_blocks.5.1.transformer_blocks.0.attn1", + 7: "input_blocks.5.1.transformer_blocks.0.attn2", + 8: "input_blocks.7.1.transformer_blocks.0.attn1", + 9: "input_blocks.7.1.transformer_blocks.0.attn2", + 10: "input_blocks.8.1.transformer_blocks.0.attn1", + 11: "input_blocks.8.1.transformer_blocks.0.attn2", + 12: "output_blocks.3.1.transformer_blocks.0.attn1", + 13: "output_blocks.3.1.transformer_blocks.0.attn2", + 14: "output_blocks.4.1.transformer_blocks.0.attn1", + 15: "output_blocks.4.1.transformer_blocks.0.attn2", + 16: "output_blocks.5.1.transformer_blocks.0.attn1", + 17: "output_blocks.5.1.transformer_blocks.0.attn2", + 18: "output_blocks.6.1.transformer_blocks.0.attn1", + 19: "output_blocks.6.1.transformer_blocks.0.attn2", + 20: "output_blocks.7.1.transformer_blocks.0.attn1", + 21: "output_blocks.7.1.transformer_blocks.0.attn2", + 22: "output_blocks.8.1.transformer_blocks.0.attn1", + 23: "output_blocks.8.1.transformer_blocks.0.attn2", + 24: "output_blocks.9.1.transformer_blocks.0.attn1", + 25: "output_blocks.9.1.transformer_blocks.0.attn2", + 26: "output_blocks.10.1.transformer_blocks.0.attn1", + 27: "output_blocks.10.1.transformer_blocks.0.attn2", + 28: "output_blocks.11.1.transformer_blocks.0.attn1", + 29: "output_blocks.11.1.transformer_blocks.0.attn2", + 30: "middle_block.1.transformer_blocks.0.attn1", + 31: "middle_block.1.transformer_blocks.0.attn2", +} + + +def compute_cond_mark(cond_or_uncond, sigmas): + cond_or_uncond_size = int(sigmas.shape[0]) + + cond_mark = [] + for cx in cond_or_uncond: + cond_mark += [cx] * cond_or_uncond_size + + cond_mark = torch.Tensor(cond_mark).to(sigmas) + return cond_mark + + +class LoRALinearLayer(torch.nn.Module): + def __init__(self, in_features: int, out_features: int, rank: int = 256, org=None): + super().__init__() + self.down = torch.nn.Linear(in_features, rank, bias=False) + self.up = torch.nn.Linear(rank, out_features, bias=False) + self.org = [org] + + def forward(self, h): + org_weight = self.org[0].weight.to(h) + org_bias = self.org[0].bias.to(h) if self.org[0].bias is not None else None + down_weight = self.down.weight + up_weight = self.up.weight + final_weight = org_weight + torch.mm(up_weight, down_weight) + return torch.nn.functional.linear(h, final_weight, org_bias) + + +class AttentionSharingUnit(torch.nn.Module): + # `transformer_options` passed to the most recent BasicTransformerBlock.forward + # call. + transformer_options: dict = {} + + def __init__(self, module, frames=2, use_control=True, rank=256): + super().__init__() + + self.heads = module.heads + self.frames = frames + self.original_module = [module] + q_in_channels, q_out_channels = ( + module.to_q.in_features, + module.to_q.out_features, + ) + k_in_channels, k_out_channels = ( + module.to_k.in_features, + module.to_k.out_features, + ) + v_in_channels, v_out_channels = ( + module.to_v.in_features, + module.to_v.out_features, + ) + o_in_channels, o_out_channels = ( + module.to_out[0].in_features, + module.to_out[0].out_features, + ) + + hidden_size = k_out_channels + + self.to_q_lora = [ + LoRALinearLayer(q_in_channels, q_out_channels, rank, module.to_q) + for _ in range(self.frames) + ] + self.to_k_lora = [ + LoRALinearLayer(k_in_channels, k_out_channels, rank, module.to_k) + for _ in range(self.frames) + ] + self.to_v_lora = [ + LoRALinearLayer(v_in_channels, v_out_channels, rank, module.to_v) + for _ in range(self.frames) + ] + self.to_out_lora = [ + LoRALinearLayer(o_in_channels, o_out_channels, rank, module.to_out[0]) + for _ in range(self.frames) + ] + + self.to_q_lora = torch.nn.ModuleList(self.to_q_lora) + self.to_k_lora = torch.nn.ModuleList(self.to_k_lora) + self.to_v_lora = torch.nn.ModuleList(self.to_v_lora) + self.to_out_lora = torch.nn.ModuleList(self.to_out_lora) + + self.temporal_i = torch.nn.Linear( + in_features=hidden_size, out_features=hidden_size + ) + self.temporal_n = torch.nn.LayerNorm( + hidden_size, elementwise_affine=True, eps=1e-6 + ) + self.temporal_q = torch.nn.Linear( + in_features=hidden_size, out_features=hidden_size + ) + self.temporal_k = torch.nn.Linear( + in_features=hidden_size, out_features=hidden_size + ) + self.temporal_v = torch.nn.Linear( + in_features=hidden_size, out_features=hidden_size + ) + self.temporal_o = torch.nn.Linear( + in_features=hidden_size, out_features=hidden_size + ) + + self.control_convs = None + + if use_control: + self.control_convs = [ + torch.nn.Sequential( + torch.nn.Conv2d(256, 256, kernel_size=3, padding=1, stride=1), + torch.nn.SiLU(), + torch.nn.Conv2d(256, hidden_size, kernel_size=1), + ) + for _ in range(self.frames) + ] + self.control_convs = torch.nn.ModuleList(self.control_convs) + + self.control_signals = None + + def forward(self, h, context=None, value=None): + transformer_options = self.transformer_options + + modified_hidden_states = einops.rearrange( + h, "(b f) d c -> f b d c", f=self.frames + ) + + if self.control_convs is not None: + context_dim = int(modified_hidden_states.shape[2]) + control_outs = [] + for f in range(self.frames): + control_signal = self.control_signals[context_dim].to( + modified_hidden_states + ) + control = self.control_convs[f](control_signal) + control = einops.rearrange(control, "b c h w -> b (h w) c") + control_outs.append(control) + control_outs = torch.stack(control_outs, dim=0) + modified_hidden_states = modified_hidden_states + control_outs.to( + modified_hidden_states + ) + + if context is None: + framed_context = modified_hidden_states + else: + framed_context = einops.rearrange( + context, "(b f) d c -> f b d c", f=self.frames + ) + + framed_cond_mark = einops.rearrange( + compute_cond_mark( + transformer_options["cond_or_uncond"], + transformer_options["sigmas"], + ), + "(b f) -> f b", + f=self.frames, + ).to(modified_hidden_states) + + attn_outs = [] + for f in range(self.frames): + fcf = framed_context[f] + + if context is not None: + cond_overwrite = transformer_options.get("cond_overwrite", []) + if len(cond_overwrite) > f: + cond_overwrite = cond_overwrite[f] + else: + cond_overwrite = None + if cond_overwrite is not None: + cond_mark = framed_cond_mark[f][:, None, None] + fcf = cond_overwrite.to(fcf) * (1.0 - cond_mark) + fcf * cond_mark + + q = self.to_q_lora[f](modified_hidden_states[f]) + k = self.to_k_lora[f](fcf) + v = self.to_v_lora[f](fcf) + o = optimized_attention(q, k, v, self.heads) + o = self.to_out_lora[f](o) + o = self.original_module[0].to_out[1](o) + attn_outs.append(o) + + attn_outs = torch.stack(attn_outs, dim=0) + modified_hidden_states = modified_hidden_states + attn_outs.to( + modified_hidden_states + ) + modified_hidden_states = einops.rearrange( + modified_hidden_states, "f b d c -> (b f) d c", f=self.frames + ) + + x = modified_hidden_states + x = self.temporal_n(x) + x = self.temporal_i(x) + d = x.shape[1] + + x = einops.rearrange(x, "(b f) d c -> (b d) f c", f=self.frames) + + q = self.temporal_q(x) + k = self.temporal_k(x) + v = self.temporal_v(x) + + x = optimized_attention(q, k, v, self.heads) + x = self.temporal_o(x) + x = einops.rearrange(x, "(b d) f c -> (b f) d c", d=d) + + modified_hidden_states = modified_hidden_states + x + + return modified_hidden_states - h + + @classmethod + def hijack_transformer_block(cls): + def register_get_transformer_options(func): + @functools.wraps(func) + def forward(self, x, context=None, transformer_options={}): + cls.transformer_options = transformer_options + return func(self, x, context, transformer_options) + + return forward + + from comfy.ldm.modules.attention import BasicTransformerBlock + + BasicTransformerBlock.forward = register_get_transformer_options( + BasicTransformerBlock.forward + ) + + +AttentionSharingUnit.hijack_transformer_block() + + +class AdditionalAttentionCondsEncoder(torch.nn.Module): + def __init__(self): + super().__init__() + + self.blocks_0 = torch.nn.Sequential( + torch.nn.Conv2d(3, 32, kernel_size=3, padding=1, stride=1), + torch.nn.SiLU(), + torch.nn.Conv2d(32, 32, kernel_size=3, padding=1, stride=1), + torch.nn.SiLU(), + torch.nn.Conv2d(32, 64, kernel_size=3, padding=1, stride=2), + torch.nn.SiLU(), + torch.nn.Conv2d(64, 64, kernel_size=3, padding=1, stride=1), + torch.nn.SiLU(), + torch.nn.Conv2d(64, 128, kernel_size=3, padding=1, stride=2), + torch.nn.SiLU(), + torch.nn.Conv2d(128, 128, kernel_size=3, padding=1, stride=1), + torch.nn.SiLU(), + torch.nn.Conv2d(128, 256, kernel_size=3, padding=1, stride=2), + torch.nn.SiLU(), + torch.nn.Conv2d(256, 256, kernel_size=3, padding=1, stride=1), + torch.nn.SiLU(), + ) # 64*64*256 + + self.blocks_1 = torch.nn.Sequential( + torch.nn.Conv2d(256, 256, kernel_size=3, padding=1, stride=2), + torch.nn.SiLU(), + torch.nn.Conv2d(256, 256, kernel_size=3, padding=1, stride=1), + torch.nn.SiLU(), + ) # 32*32*256 + + self.blocks_2 = torch.nn.Sequential( + torch.nn.Conv2d(256, 256, kernel_size=3, padding=1, stride=2), + torch.nn.SiLU(), + torch.nn.Conv2d(256, 256, kernel_size=3, padding=1, stride=1), + torch.nn.SiLU(), + ) # 16*16*256 + + self.blocks_3 = torch.nn.Sequential( + torch.nn.Conv2d(256, 256, kernel_size=3, padding=1, stride=2), + torch.nn.SiLU(), + torch.nn.Conv2d(256, 256, kernel_size=3, padding=1, stride=1), + torch.nn.SiLU(), + ) # 8*8*256 + + self.blks = [self.blocks_0, self.blocks_1, self.blocks_2, self.blocks_3] + + def __call__(self, h): + results = {} + for b in self.blks: + h = b(h) + results[int(h.shape[2]) * int(h.shape[3])] = h + return results + + +class HookerLayers(torch.nn.Module): + def __init__(self, layer_list): + super().__init__() + self.layers = torch.nn.ModuleList(layer_list) + + +class AttentionSharingPatcher(torch.nn.Module): + def __init__(self, unet, frames=2, use_control=True, rank=256): + super().__init__() + model_management.unload_model_clones(unet) + + units = [] + for i in range(32): + real_key = module_mapping_sd15[i] + attn_module = utils.get_attr(unet.model.diffusion_model, real_key) + u = AttentionSharingUnit( + attn_module, frames=frames, use_control=use_control, rank=rank + ) + units.append(u) + unet.add_object_patch("diffusion_model." + real_key, u) + + self.hookers = HookerLayers(units) + + if use_control: + self.kwargs_encoder = AdditionalAttentionCondsEncoder() + else: + self.kwargs_encoder = None + + self.dtype = torch.float32 + if model_management.should_use_fp16(model_management.get_torch_device()): + self.dtype = torch.float16 + self.hookers.half() + return + + def set_control(self, img): + img = img.cpu().float() * 2.0 - 1.0 + signals = self.kwargs_encoder(img) + for m in self.hookers.layers: + m.control_signals = signals + return \ No newline at end of file diff --git a/py/layer_diffuse/func.py b/py/layer_diffuse/func.py new file mode 100644 index 0000000..3f93d65 --- /dev/null +++ b/py/layer_diffuse/func.py @@ -0,0 +1,138 @@ +import torch +import comfy.model_management +from enum import Enum +from comfy.utils import load_torch_file +from comfy.conds import CONDRegular +from comfy_extras.nodes_compositing import JoinImageWithAlpha +from .model import ModelPatcher, TransparentVAEDecoder, calculate_weight_adjust_channel +from .attension_sharing import AttentionSharingPatcher +from ..config import LAYER_DIFFUSION, LAYER_DIFFUSION_DIR, LAYER_DIFFUSION_VAE +from ..libs.utils import to_lora_patch_dict, get_local_filepath, get_sd_version + +class LayerMethod(Enum): + FG_ONLY_ATTN = "Attention Injection" + FG_ONLY_CONV = "Conv Injection" + FG_TO_BLEND = "Foreground" + FG_BLEND_TO_BG = "Foreground to Background" + BG_TO_BLEND = "Background" + BG_BLEND_TO_FG = "Background to Foreground" + +class LayerDiffuse: + + def __init__(self) -> None: + self.vae_transparent_decoder = None + self.frames = 1 + try: + ModelPatcher.calculate_weight = calculate_weight_adjust_channel(ModelPatcher.calculate_weight) + except: + pass + + def get_layer_diffusion_method(self, method, has_blend_latent): + method = LayerMethod(method) + if method == LayerMethod.BG_TO_BLEND and has_blend_latent: + method = LayerMethod.BG_BLEND_TO_FG + elif method == LayerMethod.FG_TO_BLEND and has_blend_latent: + method = LayerMethod.FG_BLEND_TO_BG + return method + + def apply_layer_c_concat(self, cond, uncond, c_concat): + def write_c_concat(cond): + new_cond = [] + for t in cond: + n = [t[0], t[1].copy()] + if "model_conds" not in n[1]: + n[1]["model_conds"] = {} + n[1]["model_conds"]["c_concat"] = CONDRegular(c_concat) + new_cond.append(n) + return new_cond + + return (write_c_concat(cond), write_c_concat(uncond)) + + def apply_layer_diffusion(self, model: ModelPatcher, method, weight, samples, blend_samples, positive, negative, control_img=None): + sd_version = get_sd_version(model) + model_url = LAYER_DIFFUSION[method.value][sd_version]["model_url"] + if method in [LayerMethod.FG_ONLY_CONV, LayerMethod.FG_ONLY_ATTN] and sd_version == 'sd15': + self.frames = 3 + if method == LayerMethod.BG_BLEND_TO_FG and sd_version == 'sd15': + self.frames = 2 + if model_url is None: + raise Exception(f"{method.value} is not supported for {sd_version} model") + model_file = get_local_filepath(model_url, LAYER_DIFFUSION_DIR) + + layer_lora_state_dict = load_torch_file(model_file) + work_model = model.clone() + if sd_version == 'sd15': + patcher = AttentionSharingPatcher( + work_model, self.frames, use_control=control_img is not None + ) + patcher.load_state_dict(layer_lora_state_dict, strict=True) + if control_img is not None: + patcher.set_control(control_img) + else: + layer_lora_patch_dict = to_lora_patch_dict(layer_lora_state_dict) + work_model.add_patches(layer_lora_patch_dict, weight) + + # cond_contact + if method in [LayerMethod.FG_ONLY_ATTN, LayerMethod.FG_ONLY_CONV]: + samp_model = work_model + else: + if method in [LayerMethod.BG_TO_BLEND, LayerMethod.FG_TO_BLEND]: + c_concat = model.model.latent_format.process_in(samples["samples"]) + else: + c_concat = model.model.latent_format.process_in(torch.cat([samples["samples"], blend_samples["samples"]], dim=1)) + samp_model, positive, negative = (work_model,) + self.apply_layer_c_concat(positive, negative, c_concat) + + return samp_model, positive, negative + + def join_image_with_alpha(self, image, alpha): + out = image.movedim(-1, 1) + if out.shape[1] == 3: # RGB + out = torch.cat([out, torch.ones_like(out[:, :1, :, :])], dim=1) + for i in range(out.shape[0]): + out[i, 3, :, :] = alpha + return out.movedim(1, -1) + + def layer_diffusion_decode(self, layer_diffusion_method, latent, blend_samples, samp_images, model): + alpha = None + if layer_diffusion_method is not None: + method = self.get_layer_diffusion_method(layer_diffusion_method, blend_samples is not None) + print(method.value) + if method in [LayerMethod.FG_ONLY_CONV, LayerMethod.FG_ONLY_ATTN, LayerMethod.BG_BLEND_TO_FG]: + if self.vae_transparent_decoder is None: + sd_version = get_sd_version(model) + print(sd_version) + if sd_version not in ['sdxl', 'sd15']: + raise Exception(f"Only SDXL and SD1.5 model supported for Layer Diffusion") + model_url = LAYER_DIFFUSION_VAE['decode'][sd_version]["model_url"] + if model_url is None: + raise Exception(f"{method.value} is not supported for {sd_version} model") + decoder_file = get_local_filepath(model_url, LAYER_DIFFUSION_DIR) + self.vae_transparent_decoder = TransparentVAEDecoder( + load_torch_file(decoder_file), + device=comfy.model_management.get_torch_device(), + dtype=(torch.float16 if comfy.model_management.should_use_fp16() else torch.float32), + ) + + pixel = samp_images.movedim(-1, 1) # [B, H, W, C] => [B, C, H, W] + decoded = [] + sub_batch_size = 16 + for start_idx in range(0, latent.shape[0], sub_batch_size): + decoded.append( + self.vae_transparent_decoder.decode_pixel( + pixel[start_idx: start_idx + sub_batch_size], + latent[start_idx: start_idx + sub_batch_size], + ) + ) + pixel_with_alpha = torch.cat(decoded, dim=0) + # [B, C, H, W] => [B, H, W, C] + pixel_with_alpha = pixel_with_alpha.movedim(1, -1) + image = pixel_with_alpha[..., 1:] + alpha = pixel_with_alpha[..., 0] + + alpha = 1.0 - alpha + new_images, = JoinImageWithAlpha().join_image_with_alpha(image, alpha) + else: + new_images = samp_images + else: + new_images = samp_images + return (new_images, samp_images, alpha) \ No newline at end of file diff --git a/py/layer_diffusion.py b/py/layer_diffuse/model.py similarity index 74% rename from py/layer_diffusion.py rename to py/layer_diffuse/model.py index a6d362d..61744f7 100644 --- a/py/layer_diffusion.py +++ b/py/layer_diffuse/model.py @@ -5,17 +5,9 @@ import numpy as np import comfy.model_management from comfy.model_patcher import ModelPatcher -from enum import Enum from tqdm import tqdm from typing import Optional, Tuple -class LayerMethod(Enum): - FG_ONLY_ATTN = "Attention Injection" - FG_ONLY_CONV = "Conv Injection" - FG_TO_BLEND = "Foreground" - FG_BLEND_TO_BG = "Foreground to Background" - BG_TO_BLEND = "Background" - BG_BLEND_TO_FG = "Background to Foreground" try: from diffusers.configuration_utils import ConfigMixin, register_to_config @@ -383,100 +375,4 @@ except ImportError: print("\33[31mpip install diffusers\033[0m") -from comfy.utils import load_torch_file -from comfy.conds import CONDRegular -from comfy_extras.nodes_compositing import JoinImageWithAlpha -from .config import LAYER_DIFFUSION, LAYER_DIFFUSION_DIR, LAYER_DIFFUSION_VAE -from .libs.utils import to_lora_patch_dict, get_local_filepath -class LayerDiffuse: - def __init__(self) -> None: - self.vae_transparent_decoder = None - self.vae_transparent_encoder = None - - def get_layer_diffusion_method(self, method, has_blend_latent): - method = LayerMethod(method) - if method == LayerMethod.BG_TO_BLEND and has_blend_latent: - method = LayerMethod.BG_BLEND_TO_FG - elif method == LayerMethod.FG_TO_BLEND and has_blend_latent: - method = LayerMethod.FG_BLEND_TO_BG - return method - - def apply_layer_c_concat(self, cond, uncond, c_concat): - def write_c_concat(cond): - new_cond = [] - for t in cond: - n = [t[0], t[1].copy()] - if "model_conds" not in n[1]: - n[1]["model_conds"] = {} - n[1]["model_conds"]["c_concat"] = CONDRegular(c_concat) - new_cond.append(n) - return new_cond - - return (write_c_concat(cond), write_c_concat(uncond)) - - def apply_layer_diffusion(self, model: ModelPatcher, method, weight, samples, blend_samples, positive, negative): - - model_file = get_local_filepath(LAYER_DIFFUSION[method.value]["model_url"], LAYER_DIFFUSION_DIR) - layer_lora_state_dict = load_torch_file(model_file) - layer_lora_patch_dict = to_lora_patch_dict(layer_lora_state_dict) - work_model = model.clone() - work_model.add_patches(layer_lora_patch_dict, weight) - - # cond_contact - if method in [LayerMethod.FG_ONLY_ATTN, LayerMethod.FG_ONLY_CONV]: - samp_model = work_model - else: - if method in [LayerMethod.BG_TO_BLEND, LayerMethod.FG_TO_BLEND]: - c_concat = model.model.latent_format.process_in(samples["samples"]) - else: - c_concat = model.model.latent_format.process_in(torch.cat([samples["samples"], blend_samples["samples"]], dim=1)) - samp_model, positive, negative = (work_model,) + self.apply_layer_c_concat(positive, negative, c_concat) - - return samp_model, positive, negative - - def join_image_with_alpha(self, image, alpha): - out = image.movedim(-1, 1) - if out.shape[1] == 3: # RGB - out = torch.cat([out, torch.ones_like(out[:, :1, :, :])], dim=1) - for i in range(out.shape[0]): - out[i, 3, :, :] = alpha - return out.movedim(1, -1) - - def layer_diffusion_decode(self, layer_diffusion_method, latent, blend_samples, samp_images): - alpha = None - if layer_diffusion_method is not None: - method = self.get_layer_diffusion_method(layer_diffusion_method, blend_samples is not None) - print(method.value) - if method in [LayerMethod.FG_ONLY_CONV, LayerMethod.FG_ONLY_ATTN, LayerMethod.BG_BLEND_TO_FG]: - if self.vae_transparent_decoder is None: - decoder_file = get_local_filepath(LAYER_DIFFUSION_VAE['decode']["model_url"], LAYER_DIFFUSION_DIR) - self.vae_transparent_decoder = TransparentVAEDecoder( - load_torch_file(decoder_file), - device=comfy.model_management.get_torch_device(), - dtype=(torch.float16 if comfy.model_management.should_use_fp16() else torch.float32), - ) - - pixel = samp_images.movedim(-1, 1) # [B, H, W, C] => [B, C, H, W] - decoded = [] - sub_batch_size = 16 - for start_idx in range(0, latent.shape[0], sub_batch_size): - decoded.append( - self.vae_transparent_decoder.decode_pixel( - pixel[start_idx: start_idx + sub_batch_size], - latent[start_idx: start_idx + sub_batch_size], - ) - ) - pixel_with_alpha = torch.cat(decoded, dim=0) - # [B, C, H, W] => [B, H, W, C] - pixel_with_alpha = pixel_with_alpha.movedim(1, -1) - image = pixel_with_alpha[..., 1:] - alpha = pixel_with_alpha[..., 0] - - alpha = 1.0 - alpha - new_images, = JoinImageWithAlpha().join_image_with_alpha(image, alpha) - else: - new_images = samp_images - else: - new_images = samp_images - return (new_images, samp_images, alpha) \ No newline at end of file diff --git a/py/libs/utils.py b/py/libs/utils.py index 9ae54ca..f0e7740 100644 --- a/py/libs/utils.py +++ b/py/libs/utils.py @@ -27,6 +27,21 @@ def add_folder_path_and_extensions(folder_name, full_folder_paths, extensions): else: folder_paths.folder_names_and_paths[folder_name] = (full_folder_paths, extensions) +from comfy.model_base import BaseModel +import comfy.supported_models +import comfy.supported_models_base +def get_sd_version(model): + base: BaseModel = model.model + model_config: comfy.supported_models.supported_models_base.BASE = base.model_config + if isinstance(model_config, comfy.supported_models.SDXL): + return 'sdxl' + elif isinstance( + model_config, (comfy.supported_models.SD15, comfy.supported_models.SD20) + ): + return 'sd15' + else: + return 'unknown' + def find_nearest_steps(clip_id, prompt): """Find the nearest KSampler or preSampling node that references the given id.""" def check_link_to_clip(node_id, clip_id, visited=None, node=None): diff --git a/requirements.txt b/requirements.txt index 19c4786..0462bd2 100644 --- a/requirements.txt +++ b/requirements.txt @@ -1,2 +1,2 @@ -diffusers==0.25.0 +diffusers>=0.25.0 aiohttp \ No newline at end of file