fix type annotations and add within-batch node
This commit is contained in:
+54
-15
@@ -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",
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user