fix type annotations and add within-batch node

This commit is contained in:
brianf
2023-12-12 22:50:13 +00:00
parent 854a461451
commit c4f912b310
+54 -15
View File
@@ -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",
}