diff --git a/__init__.py b/__init__.py index 61b5806..31dea5d 100644 --- a/__init__.py +++ b/__init__.py @@ -51,7 +51,7 @@ def concat_first(feat: T, dim=2, scale=1.0) -> T: return torch.cat((feat, feat_style), dim=dim) -def calc_mean_std(feat, eps: float = 1e-5) -> tuple[T, T]: +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 @@ -94,7 +94,7 @@ class SharedAttentionProcessor: def get_norm_layers( layer: nn.Module, - norm_layers_: dict[str, list[Union[nn.GroupNorm, nn.LayerNorm]]], + norm_layers_: "dict[str, list[Union[nn.GroupNorm, nn.LayerNorm]]]", share_layer_norm: bool, share_group_norm: bool, ): @@ -119,7 +119,7 @@ def register_norm_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) + hidden_states = orig_forward(hidden_states) # type: ignore return hidden_states[..., :n, :] norm_layer.forward = forward_ # type: ignore @@ -141,14 +141,17 @@ def register_shared_norm( ] -class StyleAlignedSample: +SHARE_NORM_OPTIONS = ["both", "group", "layer", "disabled"] + + +class StyleAlignWithRefSampler: @classmethod def INPUT_TYPES(cls): return { "required": { "model": ("MODEL",), - "share_norm": (["both", "group", "layer", "disabled"],), - "scale": ("FLOAT", {"default": 1, "min": 0, "max": 1.0, "step": 0.1}), + "share_norm": (SHARE_NORM_OPTIONS,), + "scale": ("FLOAT", {"default": 1, "min": 0, "max": 2.0, "step": 0.1}), "batch_size": ("INT", {"default": 2, "min": 1, "max": 8, "step": 1}), "noise_seed": ( "INT", @@ -175,7 +178,7 @@ class StyleAlignedSample: RETURN_TYPES = ("LATENT", "LATENT") RETURN_NAMES = ("output", "denoised_output") FUNCTION = "patch" - CATEGORY = "advanced" + CATEGORY = "style_aligned" def __init__(self) -> None: self.args = StyleAlignedArgs() @@ -192,11 +195,10 @@ class StyleAlignedSample: negative: T, sampler: T, sigmas: T, - ref_latent: dict[str, T], - ) -> tuple[dict, dict]: + ref_latent: "dict[str, T]", + ) -> "tuple[dict, dict]": m = model.clone() - # Concat batch with style latent style_latent_tensor = ref_latent["samples"] height, width = style_latent_tensor.shape[-2:] @@ -237,7 +239,7 @@ class StyleAlignedSample: disable_pbar=disable_pbar, seed=noise_seed, ) - + # remove reference image samples = samples[1:] @@ -246,14 +248,51 @@ class StyleAlignedSample: if "x0" in x0_output: out_denoised = latent.copy() x0 = x0_output["x0"][1:] - out_denoised["samples"] = m.model.process_latent_out( - x0.cpu() - ) + out_denoised["samples"] = m.model.process_latent_out(x0.cpu()) else: out_denoised = out return (out, out_denoised) +class StyleAlignWithinBatch: + @classmethod + def INPUT_TYPES(cls): + return { + "required": { + "model": ("MODEL",), + "share_norm": (SHARE_NORM_OPTIONS,), + "scale": ("FLOAT", {"default": 1, "min": 0, "max": 1.0, "step": 0.1}), + } + } + + RETURN_TYPES = ("MODEL",) + FUNCTION = "patch" + CATEGORY = "style_aligned" + + def __init__(self) -> None: + self.args = StyleAlignedArgs() + + def patch( + self, + model: ModelPatcher, + share_norm: str, + scale: float, + ): + m = model.clone() + share_group_norm = share_norm in ["group", "both"] + share_layer_norm = share_norm in ["layer", "both"] + register_shared_norm(model, share_group_norm, share_layer_norm) + m.set_model_attn1_patch(SharedAttentionProcessor(self.args, scale)) + return (m,) + + NODE_CLASS_MAPPINGS = { - "StyleAlignedSample": StyleAlignedSample, + "StyleAlignWithRefSampler": StyleAlignWithRefSampler, + "StyleAlignWithinBatch": StyleAlignWithinBatch, +} + + +NODE_DISPLAY_NAME_MAPPINGS = { + "StyleAlignWithRefSampler": "StyleAligned Reference Sampler", + "StyleAlignWithinBatch": "StyleAligned Within Batch", }