share group and layer norms

This commit is contained in:
Brian Fitzgerald
2023-12-10 10:36:59 -06:00
parent c238be13ef
commit d40e35e0a7
+55 -42
View File
@@ -6,7 +6,7 @@ import einops
from comfy.model_patcher import ModelPatcher from comfy.model_patcher import ModelPatcher
from comfy.ldm.modules.attention import optimized_attention, optimized_attention_masked from comfy.ldm.modules.attention import optimized_attention, optimized_attention_masked
import comfy.ops import comfy.ops
from typing import Union from typing import Optional, Union
T = torch.Tensor T = torch.Tensor
@@ -77,16 +77,14 @@ def sdpa(q: T, k: T, v: T, mask=None, heads: int = 8) -> T:
class SharedAttentionProcessor: class SharedAttentionProcessor:
def __init__( def __init__(self, args: StyleAlignedArgs, style_image: Optional[T]):
self,
args: StyleAlignedArgs,
style_image: T
):
self.args = args self.args = args
self.ref_img = style_image self.ref_img = style_image
def __call__(self, q, k, v, extra_options): def __call__(self, q, k, v, extra_options):
current_index = "{}_{}".format(extra_options["transformer_index"], extra_options["block_index"]) current_index = "{}_{}".format(
extra_options["transformer_index"], extra_options["block_index"]
)
if self.args.adain_queries: if self.args.adain_queries:
q = adain(q) q = adain(q)
@@ -100,40 +98,55 @@ class SharedAttentionProcessor:
return q, k, v return q, k, v
def get_norm_layers(
layer: nn.Module,
norm_layers_: dict[str, list[Union[nn.GroupNorm, nn.LayerNorm]]],
share_layer_norm: bool,
share_group_norm: bool,
):
print(
"get layer",
layer.__class__.__name__,
isinstance(layer, nn.LayerNorm),
isinstance(layer, nn.GroupNorm),
)
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 child_layer in layer.children():
get_norm_layers(
child_layer, norm_layers_, share_layer_norm, share_group_norm
)
def register_norm_forward(
norm_layer: Union[nn.GroupNorm, nn.LayerNorm],
) -> Union[nn.GroupNorm, nn.LayerNorm]:
if not hasattr(norm_layer, "orig_forward"):
setattr(norm_layer, "orig_forward", norm_layer.forward)
orig_forward = norm_layer.orig_forward
def forward_(hidden_states: T) -> T:
n = hidden_states.shape[-2]
hidden_states = concat_first(hidden_states, dim=-2)
hidden_states = orig_forward(hidden_states)
return hidden_states[..., :n, :]
norm_layer.forward = forward_ # type: ignore
return norm_layer
def register_shared_norm( def register_shared_norm(
model: ModelPatcher, model: ModelPatcher,
share_group_norm: bool = True, share_group_norm: bool = True,
share_layer_norm: bool = True, share_layer_norm: bool = True,
): ):
def register_norm_forward(
norm_layer: Union[nn.GroupNorm, nn.LayerNorm],
) -> Union[nn.GroupNorm, nn.LayerNorm]:
if not hasattr(norm_layer, "orig_forward"):
setattr(norm_layer, "orig_forward", norm_layer.forward)
orig_forward = norm_layer.orig_forward
def forward_(hidden_states: T) -> T:
n = hidden_states.shape[-2]
hidden_states = concat_first(hidden_states, dim=-2)
hidden_states = orig_forward(hidden_states)
return hidden_states[..., :n, :]
norm_layer.forward = forward_ # type: ignore
return norm_layer
def get_norm_layers(
layer, norm_layers_: dict[str, list[Union[nn.GroupNorm, nn.LayerNorm]]]
):
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 layer.children():
get_norm_layers(layer, norm_layers_)
norm_layers = {"group": [], "layer": []} norm_layers = {"group": [], "layer": []}
get_norm_layers(model, norm_layers) get_norm_layers(model.model, norm_layers, share_layer_norm, share_group_norm)
print(f"Patching {len(norm_layers['group'])} group norms, {len(norm_layers['layer'])} layer norms.")
return [register_norm_forward(layer) for layer in norm_layers["group"]] + [ return [register_norm_forward(layer) for layer in norm_layers["group"]] + [
register_norm_forward(layer) for layer in norm_layers["layer"] register_norm_forward(layer) for layer in norm_layers["layer"]
] ]
@@ -157,15 +170,16 @@ def _get_switch_vec(total_num_layers, level):
return vec return vec
class StyleAlignedPatch: class StyleAlignedPatch:
@classmethod @classmethod
def INPUT_TYPES(cls): def INPUT_TYPES(cls):
return { return {
"required": { "required": {
"model": ("MODEL",), "model": ("MODEL",),
},
"optional": {
"style_image": ("IMAGE",), "style_image": ("IMAGE",),
} },
} }
RETURN_TYPES = ("MODEL",) RETURN_TYPES = ("MODEL",)
@@ -174,13 +188,12 @@ class StyleAlignedPatch:
def __init__(self) -> None: def __init__(self) -> None:
self.args = StyleAlignedArgs() self.args = StyleAlignedArgs()
# TODO patch norm layers
# self.norm_layers = register_shared_norm(
# model, self.args.share_group_norm, self.args.share_layer_norm
# )
def patch(self, model: ModelPatcher, style_image: T): def patch(self, model: ModelPatcher, style_image: Optional[T] = None):
m = model.clone() m = model.clone()
register_shared_norm(
model, self.args.share_group_norm, self.args.share_layer_norm
)
m.set_model_attn1_patch(SharedAttentionProcessor(self.args, style_image)) m.set_model_attn1_patch(SharedAttentionProcessor(self.args, style_image))
return (m,) return (m,)