From 6404f330c1a6314d48cefdfe1a67750dd8471797 Mon Sep 17 00:00:00 2001 From: Brian Fitzgerald Date: Fri, 8 Dec 2023 20:13:55 -0600 Subject: [PATCH] wip --- README.md | 3 + nodes.py | 238 ++++++++++++++++++++++++++++++++++++++---------------- 2 files changed, 173 insertions(+), 68 deletions(-) create mode 100644 README.md diff --git a/README.md b/README.md new file mode 100644 index 0000000..e4e494d --- /dev/null +++ b/README.md @@ -0,0 +1,3 @@ +# StyleAligned for ComfyUI + +Implementation of the [StyleAligned](https://style-aligned-gen.github.io/) paper for ComfyUI. Work in progress. \ No newline at end of file diff --git a/nodes.py b/nodes.py index e75536f..0dde56b 100644 --- a/nodes.py +++ b/nodes.py @@ -3,10 +3,22 @@ import torch import torch.nn as nn from torch.nn import functional as nnf import einops +from comfy.model_patcher import ModelPatcher +from comfy.ldm.modules.attention import optimized_attention, optimized_attention_masked T = torch.Tensor +def exists(val): + return val is not None + + +def default(val, d): + if exists(val): + return val + return d + + @dataclass(frozen=True) class StyleAlignedArgs: share_group_norm: bool = True @@ -54,39 +66,105 @@ def adain(feat: T) -> T: feat = feat * feat_style_std + feat_style_mean return feat -class SharedAttentionProcessor: - def shifted_scaled_dot_product_attention(self, attn, query: T, key: T, value: T) -> T: - logits = torch.einsum('bhqd,bhkd->bhqk', query, key) * attn.scale - logits[:, :, :, query.shape[2]:] += self.shared_score_shift +class CrossAttention(nn.Module): + def __init__( + self, + query_dim, + context_dim=None, + heads=8, + dim_head=64, + dropout=0.0, + dtype=None, + device=None, + operations=comfy.ops, + ): + super().__init__() + inner_dim = dim_head * heads + context_dim = default(context_dim, query_dim) + + self.heads = heads + self.dim_head = dim_head + + self.to_q = operations.Linear( + query_dim, inner_dim, bias=False, dtype=dtype, device=device + ) + self.to_k = operations.Linear( + context_dim, inner_dim, bias=False, dtype=dtype, device=device + ) + self.to_v = operations.Linear( + context_dim, inner_dim, bias=False, dtype=dtype, device=device + ) + + self.to_out = nn.Sequential( + operations.Linear(inner_dim, query_dim, dtype=dtype, device=device), + nn.Dropout(dropout), + ) + + def forward(self, x, context=None, value=None, mask=None): + q = self.to_q(x) + context = default(context, x) + k = self.to_k(context) + if value is not None: + v = self.to_v(value) + del value + else: + v = self.to_v(context) + + if mask is None: + out = optimized_attention(q, k, v, self.heads) + else: + out = optimized_attention_masked(q, k, v, self.heads, mask) + return self.to_out(out) + + +class SharedAttentionProcessor: + def __init__(self, style_aligned_args: StyleAlignedArgs): + super().__init__() + self.args = style_aligned_args + + def shifted_scaled_dot_product_attention( + self, attn, query: T, key: T, value: T + ) -> T: + logits = torch.einsum("bhqd,bhkd->bhqk", query, key) * attn.scale + logits[:, :, :, query.shape[2] :] += self.args.shared_score_shift probs = logits.softmax(-1) - return torch.einsum('bhqk,bhkd->bhqd', probs, value) + return torch.einsum("bhqk,bhkd->bhqd", probs, value) def shared_call( - self, - attn, - hidden_states, - encoder_hidden_states=None, - attention_mask=None, + self, + attn, + hidden_states, + encoder_hidden_states=None, + attention_mask=None, ): - residual = hidden_states input_ndim = hidden_states.ndim if input_ndim == 4: batch_size, channel, height, width = hidden_states.shape - hidden_states = hidden_states.view(batch_size, channel, height * width).transpose(1, 2) + hidden_states = hidden_states.view( + batch_size, channel, height * width + ).transpose(1, 2) batch_size, sequence_length, _ = ( - hidden_states.shape if encoder_hidden_states is None else encoder_hidden_states.shape + hidden_states.shape + if encoder_hidden_states is None + else encoder_hidden_states.shape ) if attention_mask is not None: - attention_mask = attn.prepare_attention_mask(attention_mask, sequence_length, batch_size) + attention_mask = attn.prepare_attention_mask( + attention_mask, sequence_length, batch_size + ) # scaled_dot_product_attention expects attention_mask shape to be # (batch, heads, source_length, target_length) - attention_mask = attention_mask.view(batch_size, attn.heads, -1, attention_mask.shape[-1]) + attention_mask = attention_mask.view( + batch_size, attn.heads, -1, attention_mask.shape[-1] + ) if attn.group_norm is not None: - hidden_states = attn.group_norm(hidden_states.transpose(1, 2)).transpose(1, 2) + hidden_states = attn.group_norm(hidden_states.transpose(1, 2)).transpose( + 1, 2 + ) query = attn.to_q(hidden_states) key = attn.to_k(hidden_states) @@ -98,27 +176,44 @@ class SharedAttentionProcessor: key = key.view(batch_size, -1, attn.heads, head_dim).transpose(1, 2) value = value.view(batch_size, -1, attn.heads, head_dim).transpose(1, 2) # if self.step >= self.start_inject: - if self.adain_queries: + if self.args.adain_queries: query = adain(query) - if self.adain_keys: + if self.args.adain_keys: key = adain(key) - if self.adain_values: + if self.args.adain_values: value = adain(value) - if self.share_attention: + if self.args.share_attention: key = concat_first(key, -2, scale=self.shared_score_scale) value = concat_first(value, -2) - if self.shared_score_shift != 0: - hidden_states = self.shifted_scaled_dot_product_attention(attn, query, key, value,) + if self.args.shared_score_shift != 0: + hidden_states = self.shifted_scaled_dot_product_attention( + attn, + query, + key, + value, + ) else: hidden_states = nnf.scaled_dot_product_attention( - query, key, value, attn_mask=attention_mask, dropout_p=0.0, is_causal=False + query, + key, + value, + attn_mask=attention_mask, + dropout_p=0.0, + is_causal=False, ) else: hidden_states = nnf.scaled_dot_product_attention( - query, key, value, attn_mask=attention_mask, dropout_p=0.0, is_causal=False + query, + key, + value, + attn_mask=attention_mask, + dropout_p=0.0, + is_causal=False, ) # hidden_states = adain(hidden_states) - hidden_states = hidden_states.transpose(1, 2).reshape(batch_size, -1, attn.heads * head_dim) + hidden_states = hidden_states.transpose(1, 2).reshape( + batch_size, -1, attn.heads * head_dim + ) hidden_states = hidden_states.to(query.dtype) # linear proj @@ -127,7 +222,9 @@ class SharedAttentionProcessor: hidden_states = attn.to_out[1](hidden_states) if input_ndim == 4: - hidden_states = hidden_states.transpose(-1, -2).reshape(batch_size, channel, height, width) + hidden_states = hidden_states.transpose(-1, -2).reshape( + batch_size, channel, height, width + ) if attn.residual_connection: hidden_states = hidden_states + residual @@ -135,32 +232,39 @@ class SharedAttentionProcessor: hidden_states = hidden_states / attn.rescale_output_factor return hidden_states - def __call__(self, attn, hidden_states, encoder_hidden_states=None, - attention_mask=None, **kwargs): + def __call__( + self, + attn, + hidden_states, + encoder_hidden_states=None, + attention_mask=None, + **kwargs + ): if self.full_attention_share: b, n, d = hidden_states.shape - hidden_states = einops.rearrange(hidden_states, '(k b) n d -> k (b n) d', k=2) - hidden_states = super().__call__(attn, hidden_states, encoder_hidden_states=encoder_hidden_states, - attention_mask=attention_mask, **kwargs) - hidden_states = einops.rearrange(hidden_states, 'k (b n) d -> (k b) n d', n=n) + hidden_states = einops.rearrange( + hidden_states, "(k b) n d -> k (b n) d", k=2 + ) + hidden_states = super().__call__( + attn, + hidden_states, + encoder_hidden_states=encoder_hidden_states, + attention_mask=attention_mask, + **kwargs + ) + hidden_states = einops.rearrange( + hidden_states, "k (b n) d -> (k b) n d", n=n + ) else: - hidden_states = self.shared_call(attn, hidden_states, hidden_states, attention_mask, **kwargs) + hidden_states = self.shared_call( + attn, hidden_states, hidden_states, attention_mask, **kwargs + ) return hidden_states - def __init__(self, style_aligned_args: StyleAlignedArgs): - super().__init__() - self.share_attention = style_aligned_args.share_attention - self.adain_queries = style_aligned_args.adain_queries - self.adain_keys = style_aligned_args.adain_keys - self.adain_values = style_aligned_args.adain_values - self.full_attention_share = style_aligned_args.full_attention_share - self.shared_score_scale = style_aligned_args.shared_score_scale - self.shared_score_shift = style_aligned_args.shared_score_shift - def register_shared_norm( - pipeline, + model: ModelPatcher, share_group_norm: bool = True, share_layer_norm: bool = True, ): @@ -181,28 +285,29 @@ def register_shared_norm( return norm_layer def get_norm_layers( - pipeline_, norm_layers_: dict[str, list[nn.GroupNorm | nn.LayerNorm]] + layer, norm_layers_: dict[str, list[nn.GroupNorm | nn.LayerNorm]] ): - if isinstance(pipeline_, nn.LayerNorm) and share_layer_norm: - norm_layers_["layer"].append(pipeline_) - if isinstance(pipeline_, nn.GroupNorm) and share_group_norm: - norm_layers_["group"].append(pipeline_) + if isinstance(layer, nn.LayerNorm) and share_layer_norm: + norm_layers_["layer"].append(layer) + if isinstance(layer, nn.GroupNorm) and share_group_norm: + norm_layers_["group"].append(layer) else: - for layer in pipeline_.children(): + for layer in layer.children(): get_norm_layers(layer, norm_layers_) norm_layers = {"group": [], "layer": []} - get_norm_layers(pipeline.unet, norm_layers) + get_norm_layers(model, norm_layers) return [register_norm_forward(layer) for layer in norm_layers["group"]] + [ register_norm_forward(layer) for layer in norm_layers["layer"] ] + def _get_switch_vec(total_num_layers, level): if level == 0: return torch.zeros(total_num_layers, dtype=torch.bool) if level == 1: return torch.ones(total_num_layers, dtype=torch.bool) - to_flip = level > .5 + to_flip = level > 0.5 if to_flip: level = 1 - level num_switch = int(level * total_num_layers) @@ -213,28 +318,26 @@ def _get_switch_vec(total_num_layers, level): vec = ~vec return vec -def init_attention_processors(pipeline, style_aligned_args: StyleAlignedArgs | None = None): + +def init_attention_processors(pipeline, style_aligned_args: StyleAlignedArgs): attn_procs = {} unet = pipeline.unet number_of_self, number_of_cross = 0, 0 - num_self_layers = len([name for name in unet.attn_processors.keys() if 'attn1' in name]) + num_self_layers = len( + [name for name in unet.attn_processors.keys() if "attn1" in name] + ) if style_aligned_args is None: only_self_vec = _get_switch_vec(num_self_layers, 1) else: - only_self_vec = _get_switch_vec(num_self_layers, style_aligned_args.only_self_level) + only_self_vec = _get_switch_vec( + num_self_layers, style_aligned_args.only_self_level + ) for i, name in enumerate(unet.attn_processors.keys()): - is_self_attention = 'attn1' in name + is_self_attention = "attn1" in name if is_self_attention: number_of_self += 1 - if style_aligned_args is None or only_self_vec[i // 2]: - attn_procs[name] = DefaultAttentionProcessor() - else: + if only_self_vec[i // 2]: attn_procs[name] = SharedAttentionProcessor(style_aligned_args) - else: - number_of_cross += 1 - attn_procs[name] = DefaultAttentionProcessor() - - unet.set_attn_processor(attn_procs) class StyleAlignedPatch: @@ -251,15 +354,14 @@ class StyleAlignedPatch: FUNCTION = "patch" CATEGORY = "custom_node_experiments" - def __init__(self) -> None: + def __init__(self, model: ModelPatcher) -> None: self.args = StyleAlignedArgs() self.norm_layers = register_shared_norm( - None, self.args.share_group_norm, self.args.share_layer_norm + model, self.args.share_group_norm, self.args.share_layer_norm ) - def patch(self, model, style_image): + def patch(self, model): m = model.clone() - sd = model.model_state_dict() return (m,)