diff --git a/README.md b/README.md index b755b5f..172dde8 100644 --- a/README.md +++ b/README.md @@ -125,6 +125,12 @@ https://github.com/shadowcz007/comfyui-mixlab-nodes/assets/12645064/e7e77f90-e43 > Conveniently load images from a fixed address on the internet to ensure that default images in the workflow can be executed. +## Style +> Apply VisualStyle Prompting , Modified from [ComfyUI_VisualStylePrompting](https://github.com/ExponentialML/ComfyUI_VisualStylePrompting) + +![](./assets/VisualStylePrompting.png) + + ## Utils > The Color node provides a color picker for easy color selection, the Font node offers built-in font selection for use with TextImage to generate text images, and the DynamicDelayByText node allows delayed execution based on the length of the input text. diff --git a/__init__.py b/__init__.py index 921375f..70b3ad3 100644 --- a/__init__.py +++ b/__init__.py @@ -601,6 +601,7 @@ from .nodes.Audio import GamePal,SpeechRecognition,SpeechSynthesis from .nodes.Utils import CreateLoraNames,CreateSampler_names,CreateCkptNames,CreateSeedNode,TESTNODE_,TESTNODE_TOKEN,AppInfo,IntNumber,FloatSlider,TextInput,ColorInput,FontInput,TextToNumber,DynamicDelayProcessor,LimitNumber,SwitchByIndex,MultiplicationNode from .nodes.Mask import OutlineMask,FeatheredMask +from .nodes.Style import ApplyVisualStylePrompting # 要导出的所有节点及其名称的字典 # 注意:名称应全局唯一 @@ -665,7 +666,8 @@ NODE_CLASS_MAPPINGS = { "Seed_":CreateSeedNode, "CkptNames_":CreateCkptNames, "SamplerNames_":CreateSampler_names, - "LoraNames_":CreateLoraNames + "LoraNames_":CreateLoraNames, + "ApplyVisualStylePrompting_":ApplyVisualStylePrompting # "LaMaInpainting":LaMaInpainting # "GamePal":GamePal } @@ -693,7 +695,8 @@ NODE_DISPLAY_NAME_MAPPINGS = { "ChinesePrompt_Mix":"ChinesePrompt ♾️Mixlab", "GamePal":"GamePal ♾️Mixlab", "RembgNode_Mix":"Removebg", - "LoraNames_":"LoraName_TriggerWords.safetensors" + "LoraNames_":"LoraName", + "ApplyVisualStylePrompting_":"Apply VisualStyle Prompting" } # web ui的节点功能 diff --git a/assets/VisualStylePrompting.png b/assets/VisualStylePrompting.png new file mode 100644 index 0000000..5a1e6a7 Binary files /dev/null and b/assets/VisualStylePrompting.png differ diff --git a/nodes/ChatGPT.py b/nodes/ChatGPT.py index f583127..6bcb52a 100644 --- a/nodes/ChatGPT.py +++ b/nodes/ChatGPT.py @@ -119,7 +119,7 @@ class ChatGPTNode: RETURN_TYPES = ("STRING","STRING","STRING",) RETURN_NAMES = ("text","messages","session_history",) FUNCTION = "generate_contextual_text" - CATEGORY = "♾️Mixlab/GPT" + CATEGORY = "♾️Mixlab/Prompt/GPT" INPUT_IS_LIST = False OUTPUT_IS_LIST = (False,False,False,) @@ -209,7 +209,7 @@ class ShowTextForGPT: OUTPUT_NODE = True OUTPUT_IS_LIST = (True,) - CATEGORY = "♾️Mixlab/GPT" + CATEGORY = "♾️Mixlab/Prompt/GPT" def run(self, text,output_dir=[""]): @@ -293,7 +293,7 @@ class CharacterInText: # OUTPUT_NODE = True OUTPUT_IS_LIST = (False,) - CATEGORY = "♾️Mixlab/GPT" + CATEGORY = "♾️Mixlab/Prompt/GPT" def run(self, text,character,start_index): # print(text,character,start_index) @@ -338,7 +338,7 @@ class TextSplitByDelimiter: # OUTPUT_NODE = True OUTPUT_IS_LIST = (True,) - CATEGORY = "♾️Mixlab/GPT" + CATEGORY = "♾️Mixlab/Prompt/GPT" def run(self, text,delimiter,start_index,skip_every,max_count): arr=[] diff --git a/nodes/Style.py b/nodes/Style.py new file mode 100644 index 0000000..574ebd2 --- /dev/null +++ b/nodes/Style.py @@ -0,0 +1,76 @@ +import comfy +import torch + +from .VisualStylePrompting.attention_functions import VisualStyleProcessor + +class ApplyVisualStylePrompting: + + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "reference_image": ("IMAGE",), + "reference_image_text": ("STRING", {"multiline": True}), + "model": ("MODEL",), + "clip": ("CLIP", ), + "vae": ("VAE", ), + "positive": ("CONDITIONING",), + "negative": ("CONDITIONING", ), + "enabled": ("BOOLEAN", {"default": True}), + "denoise": ("FLOAT", {"default": 1., "min": 0., "max": 1., "step": 1e-2}), + "batch_size": ("INT", {"default": 1, "min": 1, "max": 4096,"step":2}) + } + } + + RETURN_TYPES = ("MODEL", "CONDITIONING","CONDITIONING", "LATENT") + RETURN_NAMES = ("model", "positive", "negative", "latents") + + CATEGORY = "♾️Mixlab/Style" + + FUNCTION = "run" + + def run( + self, + reference_image, + reference_image_text, + model: comfy.model_patcher.ModelPatcher, + clip, + vae, + positive, + negative, + enabled, + denoise, + batch_size=1 + ): + + tokens = clip.tokenize(reference_image_text) + cond, pooled = clip.encode_from_tokens(tokens, return_pooled=True) + reference_image_prompt=[[cond, {"pooled_output": pooled}]] + + reference_image = reference_image.repeat(((batch_size+1)//2, 1,1,1)) + + self.model = model + reference_latent = vae.encode(reference_image[:,:,:,:3]) + + for n, m in model.model.diffusion_model.named_modules(): + if m.__class__.__name__ == "CrossAttention": + processor = VisualStyleProcessor(m, enabled=enabled) + setattr(m, 'forward', processor.visual_style_forward) + + conditioning_prompt = reference_image_prompt + positive + negative_prompt = negative * 2 + + latents = torch.zeros_like(reference_latent) + latents = torch.cat([latents] * 2) + + if denoise < 1.0: + latents[::1] = reference_latent[:1] + else: + latents[::2] = reference_latent + + denoise_mask = torch.ones_like(latents)[:, :1, ...] * denoise + + denoise_mask[0] = 0. + + return (model, conditioning_prompt, negative_prompt, {"samples": latents, "noise_mask": denoise_mask}) + diff --git a/nodes/VisualStylePrompting/attention_functions.py b/nodes/VisualStylePrompting/attention_functions.py new file mode 100644 index 0000000..29be745 --- /dev/null +++ b/nodes/VisualStylePrompting/attention_functions.py @@ -0,0 +1,45 @@ +from comfy.ldm.modules.attention import default, optimized_attention, optimized_attention_masked +from .style_functions import adain, concat_first + +class VisualStyleProcessor(object): + def __init__(self, + module_self, + keys_scale: float = 1.0, + enabled: bool = True, + adain_queries: bool = True, + adain_keys: bool = True, + adain_values: bool = False + ): + self.module_self = module_self + self.keys_scale = keys_scale + self.enabled = enabled + self.adain_queries = adain_queries + self.adain_keys = adain_keys + self.adain_values = adain_values + + def visual_style_forward(self, x, context, value, mask=None): + q = self.module_self.to_q(x) + context = default(context, x) + k = self.module_self.to_k(context) + if value is not None: + v = self.module_self.to_v(value) + del value + else: + v = self.module_self.to_v(context) + + if self.enabled: + if self.adain_queries: + q = adain(q) + if self.adain_keys: + k = adain(k) + if self.adain_values: + v = adain(v) + + k = concat_first(k, -2, self.keys_scale) + v = concat_first(v, -2) + + if mask is None: + out = optimized_attention(q, k, v, self.module_self.heads) + else: + out = optimized_attention_masked(q, k, v, self.module_self.heads, mask) + return self.module_self.to_out(out) \ No newline at end of file diff --git a/nodes/VisualStylePrompting/style_functions.py b/nodes/VisualStylePrompting/style_functions.py new file mode 100644 index 0000000..a41321e --- /dev/null +++ b/nodes/VisualStylePrompting/style_functions.py @@ -0,0 +1,60 @@ +import torch + +from einops import rearrange +from dataclasses import dataclass + +T = torch.Tensor + +@dataclass(frozen=True) +class StyleAlignedArgs: + share_group_norm: bool = True + share_layer_norm: bool = True, + share_attention: bool = True + adain_queries: bool = True + adain_keys: bool = True + adain_values: bool = False + full_attention_share: bool = False + keys_scale: float = 1. + only_self_level: float = 0. + +def expand_first(feat: T, scale=1., ) -> T: + b = feat.shape[0] + feat_style = torch.stack((feat[0], feat[b // 2])).unsqueeze(1) + if scale == 1: + feat_style = feat_style.expand(2, b // 2, *feat.shape[1:]) + else: + feat_style = feat_style.repeat(1, b // 2, 1, 1, 1) + feat_style = torch.cat([feat_style[:, :1], scale * feat_style[:, 1:]], dim=1) + return feat_style.reshape(*feat.shape) + + +def concat_first(feat: T, dim=2, scale=1.) -> T: + feat_style = expand_first(feat, scale=scale) + return torch.cat((feat, feat_style), dim=dim) + + +def calc_mean_std(feat, eps: float = 1e-5) -> tuple[T, T]: + feat_std = (feat.var(dim=-2, keepdims=True) + eps).sqrt() + feat_mean = feat.mean(dim=-2, keepdims=True) + return feat_mean, feat_std + + +def adain(feat: T) -> T: + feat_mean, feat_std = calc_mean_std(feat) + feat_style_mean = expand_first(feat_mean) + feat_style_std = expand_first(feat_std) + feat = (feat - feat_mean) / feat_std + feat = feat * feat_style_std + feat_style_mean + return feat + +def swapping_attention(key, value, chunk_size=2): + chunk_length = key.size()[0] // chunk_size # [text-condition, null-condition] + reference_image_index = [0] * chunk_length # [0 0 0 0 0] + key = rearrange(key, "(b f) d c -> b f d c", f=chunk_length) + key = key[:, reference_image_index] # ref to all + key = rearrange(key, "b f d c -> (b f) d c") + value = rearrange(value, "(b f) d c -> b f d c", f=chunk_length) + value = value[:, reference_image_index] # ref to all + value = rearrange(value, "b f d c -> (b f) d c") + + return key, value \ No newline at end of file