Use standard Apply nodes for inpainting
This commit is contained in:
@@ -1,7 +1,7 @@
|
||||
from comfy_api.latest import ComfyExtension, io
|
||||
|
||||
from .nodes_main import (ControlNetLoaderAdvanced, DiffControlNetLoaderAdvanced, AnimaLLLiteLoaderAdvanced,
|
||||
AdvancedControlNetApply, AdvancedControlNetInpaintingApply, AdvancedControlNetApplySingle)
|
||||
AdvancedControlNetApply, AdvancedControlNetApplySingle)
|
||||
from .nodes_weight import (DefaultWeights, ScaledSoftMaskedUniversalWeights, ScaledSoftUniversalWeights,
|
||||
SoftControlNetWeightsSD15, CustomControlNetWeightsSD15, CustomControlNetWeightsFlux,
|
||||
CustomControlNetWeightsAnima, SoftT2IAdapterWeights, CustomT2IAdapterWeights, ExtrasMiddleMultNode,
|
||||
@@ -32,7 +32,6 @@ class AdvancedControlNetExtension(ComfyExtension):
|
||||
LatentKeyframeBatchedGroupNode,
|
||||
LatentKeyframeGroupNode,
|
||||
AdvancedControlNetApply,
|
||||
AdvancedControlNetInpaintingApply,
|
||||
AdvancedControlNetApplySingle,
|
||||
ControlNetLoaderAdvanced,
|
||||
DiffControlNetLoaderAdvanced,
|
||||
|
||||
@@ -266,12 +266,12 @@ class AdvancedControlNetApplyDEPR(io.ComfyNode):
|
||||
io.Float.Input('strength', default=1.0, max=10.0, min=0.0, step=0.01),
|
||||
io.Float.Input('start_percent', default=0.0, max=1.0, min=0.0, step=0.001),
|
||||
io.Float.Input('end_percent', default=1.0, max=1.0, min=0.0, step=0.001),
|
||||
io.Mask.Input('mask_optional', optional=True),
|
||||
io.Mask.Input('mask_optional', display_name='effect_mask', optional=True),
|
||||
io.Custom('TIMESTEP_KEYFRAME').Input('timestep_kf', optional=True),
|
||||
io.Custom('LATENT_KEYFRAME').Input('latent_kf_override', optional=True),
|
||||
io.Custom('CONTROL_NET_WEIGHTS').Input('weights_override', optional=True),
|
||||
io.Model.Input('model_optional', optional=True),
|
||||
io.Vae.Input('vae_optional', optional=True)
|
||||
io.Model.Input('model_optional', display_name='model', optional=True),
|
||||
io.Vae.Input('vae_optional', display_name='vae', optional=True)
|
||||
],
|
||||
outputs=[
|
||||
io.Conditioning.Output('positive', is_output_list=False),
|
||||
@@ -306,12 +306,12 @@ class AdvancedControlNetApplySingleDEPR(io.ComfyNode):
|
||||
io.Float.Input('strength', default=1.0, max=10.0, min=0.0, step=0.01),
|
||||
io.Float.Input('start_percent', default=0.0, max=1.0, min=0.0, step=0.001),
|
||||
io.Float.Input('end_percent', default=1.0, max=1.0, min=0.0, step=0.001),
|
||||
io.Mask.Input('mask_optional', optional=True),
|
||||
io.Mask.Input('mask_optional', display_name='effect_mask', optional=True),
|
||||
io.Custom('TIMESTEP_KEYFRAME').Input('timestep_kf', optional=True),
|
||||
io.Custom('LATENT_KEYFRAME').Input('latent_kf_override', optional=True),
|
||||
io.Custom('CONTROL_NET_WEIGHTS').Input('weights_override', optional=True),
|
||||
io.Model.Input('model_optional', optional=True),
|
||||
io.Vae.Input('vae_optional', optional=True)
|
||||
io.Model.Input('model_optional', display_name='model', optional=True),
|
||||
io.Vae.Input('vae_optional', display_name='vae', optional=True)
|
||||
],
|
||||
outputs=[
|
||||
io.Conditioning.Output('CONDITIONING', is_output_list=False),
|
||||
@@ -341,7 +341,7 @@ class ControlNetLoaderAdvancedDEPR(io.ComfyNode):
|
||||
category='',
|
||||
inputs=[
|
||||
io.Combo.Input('control_net_name', options=folder_paths.get_filename_list("controlnet")),
|
||||
io.Custom('TIMESTEP_KEYFRAME').Input('tk_optional', optional=True)
|
||||
io.Custom('TIMESTEP_KEYFRAME').Input('tk_optional', display_name='timestep_kf', optional=True)
|
||||
],
|
||||
outputs=[
|
||||
io.ControlNet.Output('CONTROL_NET', is_output_list=False)
|
||||
@@ -371,7 +371,7 @@ class DiffControlNetLoaderAdvancedDEPR(io.ComfyNode):
|
||||
inputs=[
|
||||
io.Model.Input('model'),
|
||||
io.Combo.Input('control_net_name', options=folder_paths.get_filename_list("controlnet")),
|
||||
io.Custom('TIMESTEP_KEYFRAME').Input('tk_optional', optional=True)
|
||||
io.Custom('TIMESTEP_KEYFRAME').Input('tk_optional', display_name='timestep_kf', optional=True)
|
||||
],
|
||||
outputs=[
|
||||
io.ControlNet.Output('CONTROL_NET', is_output_list=False)
|
||||
|
||||
@@ -25,7 +25,7 @@ class TimestepKeyframeNode(io.ComfyNode):
|
||||
io.Float.Input('null_latent_kf_strength', optional=True, default=0.0, max=10.0, min=0.0, step=0.001),
|
||||
io.Boolean.Input('inherit_missing', optional=True, default=True),
|
||||
io.Int.Input('guarantee_steps', optional=True, default=1, max=9007199254740991, min=0),
|
||||
io.Mask.Input('mask_optional', optional=True)
|
||||
io.Mask.Input('mask_optional', display_name='mask', optional=True)
|
||||
],
|
||||
outputs=[
|
||||
io.Custom('TIMESTEP_KEYFRAME').Output('TIMESTEP_KF', is_output_list=False)
|
||||
@@ -80,7 +80,7 @@ class TimestepKeyframeInterpolationNode(io.ComfyNode):
|
||||
io.Custom('LATENT_KEYFRAME').Input('latent_keyframe', optional=True),
|
||||
io.Float.Input('null_latent_kf_strength', optional=True, default=0.0, max=10.0, min=0.0, step=0.001),
|
||||
io.Boolean.Input('inherit_missing', optional=True, default=True),
|
||||
io.Mask.Input('mask_optional', optional=True),
|
||||
io.Mask.Input('mask_optional', display_name='mask', optional=True),
|
||||
io.Boolean.Input('print_keyframes', optional=True, default=False)
|
||||
],
|
||||
outputs=[
|
||||
@@ -137,7 +137,7 @@ class TimestepKeyframeFromStrengthListNode(io.ComfyNode):
|
||||
io.Custom('LATENT_KEYFRAME').Input('latent_keyframe', optional=True),
|
||||
io.Float.Input('null_latent_kf_strength', optional=True, default=0.0, max=10.0, min=0.0, step=0.001),
|
||||
io.Boolean.Input('inherit_missing', optional=True, default=True),
|
||||
io.Mask.Input('mask_optional', optional=True),
|
||||
io.Mask.Input('mask_optional', display_name='mask', optional=True),
|
||||
io.Boolean.Input('print_keyframes', optional=True, default=False)
|
||||
],
|
||||
outputs=[
|
||||
@@ -226,7 +226,7 @@ class LatentKeyframeGroupNode(io.ComfyNode):
|
||||
inputs=[
|
||||
io.String.Input('index_strengths', default='', multiline=True),
|
||||
io.Custom('LATENT_KEYFRAME').Input('prev_latent_kf', optional=True),
|
||||
io.Latent.Input('latent_optional', optional=True),
|
||||
io.Latent.Input('latent_optional', display_name='latent', optional=True),
|
||||
io.Boolean.Input('print_keyframes', optional=True, default=False)
|
||||
],
|
||||
outputs=[
|
||||
|
||||
+21
-69
@@ -17,7 +17,7 @@ class ControlNetLoaderAdvanced(io.ComfyNode):
|
||||
category='Adv-ControlNet 🛂🅐🅒🅝',
|
||||
inputs=[
|
||||
io.Combo.Input('cnet', options=folder_paths.get_filename_list("controlnet")),
|
||||
io.Custom('TIMESTEP_KEYFRAME').Input('_tk_opt', optional=True)
|
||||
io.Custom('TIMESTEP_KEYFRAME').Input('_tk_opt', display_name='timestep_kf', optional=True)
|
||||
],
|
||||
outputs=[
|
||||
io.ControlNet.Output('CONTROL_NET', is_output_list=False)
|
||||
@@ -43,7 +43,7 @@ class DiffControlNetLoaderAdvanced(io.ComfyNode):
|
||||
inputs=[
|
||||
io.Model.Input('model'),
|
||||
io.Combo.Input('cnet', options=folder_paths.get_filename_list("controlnet")),
|
||||
io.Custom('TIMESTEP_KEYFRAME').Input('_tk_opt', optional=True)
|
||||
io.Custom('TIMESTEP_KEYFRAME').Input('_tk_opt', display_name='timestep_kf', optional=True)
|
||||
],
|
||||
outputs=[
|
||||
io.ControlNet.Output('CONTROL_NET', is_output_list=False)
|
||||
@@ -97,11 +97,12 @@ class AdvancedControlNetApply(io.ComfyNode):
|
||||
io.Float.Input('strength', default=1.0, max=10.0, min=0.0, step=0.01),
|
||||
io.Float.Input('start_percent', default=0.0, max=1.0, min=0.0, step=0.001),
|
||||
io.Float.Input('end_percent', default=1.0, max=1.0, min=0.0, step=0.001),
|
||||
io.Mask.Input('mask_optional', optional=True),
|
||||
io.Mask.Input('mask_optional', display_name='effect_mask', optional=True),
|
||||
io.Custom('TIMESTEP_KEYFRAME').Input('timestep_kf', optional=True),
|
||||
io.Custom('LATENT_KEYFRAME').Input('latent_kf_override', optional=True),
|
||||
io.Custom('CONTROL_NET_WEIGHTS').Input('weights_override', optional=True),
|
||||
io.Vae.Input('vae_optional', optional=True)
|
||||
io.Vae.Input('vae_optional', display_name='vae', optional=True),
|
||||
io.Mask.Input('inpaint_mask', optional=True)
|
||||
],
|
||||
outputs=[
|
||||
io.Conditioning.Output('positive', is_output_list=False),
|
||||
@@ -113,11 +114,17 @@ class AdvancedControlNetApply(io.ComfyNode):
|
||||
def execute(cls, positive, negative, control_net, image, strength, start_percent, end_percent,
|
||||
mask_optional: Tensor=None, vae_optional=None,
|
||||
timestep_kf: TimestepKeyframeGroup=None, latent_kf_override: LatentKeyframeGroup=None,
|
||||
weights_override: ControlWeights=None, control_apply_to_uncond=False, extra_concat=None):
|
||||
weights_override: ControlWeights=None, control_apply_to_uncond=False,
|
||||
inpaint_mask: Tensor=None):
|
||||
if strength == 0 or (mask_optional is not None and mask_optional.count_nonzero().item() == 0):
|
||||
return io.NodeOutput(positive, negative)
|
||||
if extra_concat is None:
|
||||
extra_concat = []
|
||||
|
||||
extra_concat = []
|
||||
if inpaint_mask is not None and getattr(control_net, "concat_mask", False):
|
||||
source_mask = 1.0 - inpaint_mask.reshape((-1, 1, inpaint_mask.shape[-2], inpaint_mask.shape[-1]))
|
||||
mask_apply = comfy.utils.common_upscale(source_mask, image.shape[2], image.shape[1], "bilinear", "center").round()
|
||||
image = image * mask_apply.movedim(1, -1).repeat(1, 1, 1, image.shape[3])
|
||||
extra_concat = [source_mask]
|
||||
|
||||
control_hint = image.movedim(-1,1)
|
||||
cnets = {}
|
||||
@@ -156,9 +163,9 @@ class AdvancedControlNetApply(io.ComfyNode):
|
||||
elif not vae_optional:
|
||||
# make sure SD3 ControlNet will get a special message instead of generic type mention
|
||||
if is_sd3_advanced_controlnet(c_net):
|
||||
raise Exception(f"SD3 ControlNet requires vae_optional input, but got None.")
|
||||
raise Exception(f"SD3 ControlNet requires vae input, but got None.")
|
||||
else:
|
||||
raise Exception(f"Type '{type(c_net).__name__}' requires vae_optional input, but got None.")
|
||||
raise Exception(f"Type '{type(c_net).__name__}' requires vae input, but got None.")
|
||||
# apply optional parameters and overrides, if provided
|
||||
if timestep_kf is not None:
|
||||
c_net.set_timestep_keyframes(timestep_kf)
|
||||
@@ -186,62 +193,6 @@ class AdvancedControlNetApply(io.ComfyNode):
|
||||
return io.NodeOutput(out[0], out[1])
|
||||
|
||||
|
||||
class AdvancedControlNetInpaintingApply(io.ComfyNode):
|
||||
@classmethod
|
||||
def define_schema(cls) -> io.Schema:
|
||||
return io.Schema(
|
||||
node_id='ACN_AdvancedControlNetInpaintingApply',
|
||||
display_name='Apply Advanced ControlNet Inpainting 🛂🅐🅒🅝',
|
||||
category='Adv-ControlNet 🛂🅐🅒🅝',
|
||||
inputs=[
|
||||
io.Conditioning.Input('positive'),
|
||||
io.Conditioning.Input('negative'),
|
||||
io.ControlNet.Input('control_net'),
|
||||
io.Vae.Input('vae'),
|
||||
io.Image.Input('image'),
|
||||
io.Mask.Input('inpaint_mask'),
|
||||
io.Float.Input('strength', default=1.0, max=10.0, min=0.0, step=0.01),
|
||||
io.Float.Input('start_percent', default=0.0, max=1.0, min=0.0, step=0.001),
|
||||
io.Float.Input('end_percent', default=1.0, max=1.0, min=0.0, step=0.001),
|
||||
io.Mask.Input('effect_mask_optional', optional=True),
|
||||
io.Custom('TIMESTEP_KEYFRAME').Input('timestep_kf', optional=True),
|
||||
io.Custom('LATENT_KEYFRAME').Input('latent_kf_override', optional=True),
|
||||
io.Custom('CONTROL_NET_WEIGHTS').Input('weights_override', optional=True)
|
||||
],
|
||||
outputs=[
|
||||
io.Conditioning.Output('positive', is_output_list=False),
|
||||
io.Conditioning.Output('negative', is_output_list=False)
|
||||
]
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def execute(cls, positive, negative, control_net, vae, image, inpaint_mask, strength, start_percent, end_percent,
|
||||
effect_mask_optional: Tensor=None, timestep_kf: TimestepKeyframeGroup=None,
|
||||
latent_kf_override: LatentKeyframeGroup=None, weights_override: ControlWeights=None):
|
||||
if not getattr(control_net, "concat_mask", False):
|
||||
raise ValueError("The provided ControlNet does not use an inpaint source mask; use Apply Advanced ControlNet instead.")
|
||||
|
||||
source_mask = 1.0 - inpaint_mask.reshape((-1, 1, inpaint_mask.shape[-2], inpaint_mask.shape[-1]))
|
||||
mask_apply = comfy.utils.common_upscale(source_mask, image.shape[2], image.shape[1], "bilinear", "center").round()
|
||||
image = image * mask_apply.movedim(1, -1).repeat(1, 1, 1, image.shape[3])
|
||||
|
||||
return AdvancedControlNetApply.execute(
|
||||
positive=positive,
|
||||
negative=negative,
|
||||
control_net=control_net,
|
||||
image=image,
|
||||
strength=strength,
|
||||
start_percent=start_percent,
|
||||
end_percent=end_percent,
|
||||
mask_optional=effect_mask_optional,
|
||||
vae_optional=vae,
|
||||
timestep_kf=timestep_kf,
|
||||
latent_kf_override=latent_kf_override,
|
||||
weights_override=weights_override,
|
||||
extra_concat=[source_mask]
|
||||
)
|
||||
|
||||
|
||||
class AdvancedControlNetApplySingle(io.ComfyNode):
|
||||
@classmethod
|
||||
def define_schema(cls) -> io.Schema:
|
||||
@@ -256,11 +207,12 @@ class AdvancedControlNetApplySingle(io.ComfyNode):
|
||||
io.Float.Input('strength', default=1.0, max=10.0, min=0.0, step=0.01),
|
||||
io.Float.Input('start_percent', default=0.0, max=1.0, min=0.0, step=0.001),
|
||||
io.Float.Input('end_percent', default=1.0, max=1.0, min=0.0, step=0.001),
|
||||
io.Mask.Input('mask_optional', optional=True),
|
||||
io.Mask.Input('mask_optional', display_name='effect_mask', optional=True),
|
||||
io.Custom('TIMESTEP_KEYFRAME').Input('timestep_kf', optional=True),
|
||||
io.Custom('LATENT_KEYFRAME').Input('latent_kf_override', optional=True),
|
||||
io.Custom('CONTROL_NET_WEIGHTS').Input('weights_override', optional=True),
|
||||
io.Vae.Input('vae_optional', optional=True)
|
||||
io.Vae.Input('vae_optional', display_name='vae', optional=True),
|
||||
io.Mask.Input('inpaint_mask', optional=True)
|
||||
],
|
||||
outputs=[
|
||||
io.Conditioning.Output('CONDITIONING', is_output_list=False),
|
||||
@@ -272,10 +224,10 @@ class AdvancedControlNetApplySingle(io.ComfyNode):
|
||||
def execute(cls, conditioning, control_net, image, strength, start_percent, end_percent,
|
||||
mask_optional: Tensor=None, vae_optional=None,
|
||||
timestep_kf: TimestepKeyframeGroup=None, latent_kf_override: LatentKeyframeGroup=None,
|
||||
weights_override: ControlWeights=None):
|
||||
weights_override: ControlWeights=None, inpaint_mask: Tensor=None):
|
||||
values = AdvancedControlNetApply.execute(positive=conditioning, negative=None, control_net=control_net, image=image,
|
||||
strength=strength, start_percent=start_percent, end_percent=end_percent,
|
||||
mask_optional=mask_optional, vae_optional=vae_optional,
|
||||
timestep_kf=timestep_kf, latent_kf_override=latent_kf_override, weights_override=weights_override,
|
||||
control_apply_to_uncond=True)
|
||||
control_apply_to_uncond=True, inpaint_mask=inpaint_mask)
|
||||
return io.NodeOutput(values.args[0], None)
|
||||
|
||||
@@ -24,7 +24,7 @@ class SparseCtrlLoaderAdvanced(io.ComfyNode):
|
||||
io.Float.Input('motion_strength', default=1.0, max=10.0, min=0.0, step=0.001),
|
||||
io.Float.Input('motion_scale', default=1.0, max=10.0, min=0.0, step=0.001),
|
||||
io.Custom('SPARSE_METHOD').Input('sparse_method', optional=True),
|
||||
io.Custom('TIMESTEP_KEYFRAME').Input('tk_optional', optional=True),
|
||||
io.Custom('TIMESTEP_KEYFRAME').Input('tk_optional', display_name='timestep_kf', optional=True),
|
||||
io.Combo.Input('context_aware', optional=True, options=['nearest_hint', 'off']),
|
||||
io.Float.Input('sparse_hint_mult', optional=True, default=1.0, max=10.0, min=0.0, step=0.001),
|
||||
io.Float.Input('sparse_nonhint_mult', optional=True, default=1.0, max=10.0, min=0.0, step=0.001),
|
||||
@@ -60,7 +60,7 @@ class SparseCtrlMergedLoaderAdvanced(io.ComfyNode):
|
||||
io.Float.Input('motion_strength', default=1.0, max=10.0, min=0.0, step=0.001),
|
||||
io.Float.Input('motion_scale', default=1.0, max=10.0, min=0.0, step=0.001),
|
||||
io.Custom('SPARSE_METHOD').Input('sparse_method', optional=True),
|
||||
io.Custom('TIMESTEP_KEYFRAME').Input('tk_optional', optional=True)
|
||||
io.Custom('TIMESTEP_KEYFRAME').Input('tk_optional', display_name='timestep_kf', optional=True)
|
||||
],
|
||||
outputs=[
|
||||
io.ControlNet.Output('CONTROL_NET', is_output_list=False)
|
||||
|
||||
@@ -2,9 +2,9 @@
|
||||
|
||||
This reviewer example adapts the active inpainting branch of ComfyUI's official
|
||||
Qwen Image workflow. It uses **Load Advanced ControlNet Model** and **Apply
|
||||
Advanced ControlNet Inpainting**, while retaining the official Qwen base
|
||||
pipeline and the bypassed optional Lightning LoRA. Node titles are left at
|
||||
their ComfyUI defaults; the workflow stores no node title overrides.
|
||||
Advanced ControlNet**, while retaining the official Qwen base pipeline and the
|
||||
bypassed optional Lightning LoRA. Node titles are left at their ComfyUI
|
||||
defaults; the workflow stores no node title overrides.
|
||||
|
||||
## Inputs and models
|
||||
|
||||
@@ -49,10 +49,10 @@ it requires changing the sampler settings appropriately.
|
||||
The two native **Load Image** nodes are intentionally separate. **Image To
|
||||
Mask** reads the red channel of the mask PNG. That source `inpaint_mask` defines
|
||||
the region supplied to the inpainting ControlNet and the latent noise mask. It
|
||||
is not the Advanced-ControlNet effect mask. `effect_mask_optional` is left
|
||||
unconnected and independently limits where control is injected. The Apply node
|
||||
also exposes unconnected timestep keyframe, latent keyframe, and weights ports
|
||||
for focused reviewer experiments.
|
||||
is not the Advanced-ControlNet effect mask. `effect_mask` is left unconnected
|
||||
and independently limits where control is injected. The Apply node also exposes
|
||||
unconnected timestep keyframe, latent keyframe, and weights ports for focused
|
||||
reviewer experiments.
|
||||
|
||||
## Measured validation evidence
|
||||
|
||||
@@ -75,10 +75,9 @@ here rather than inferred from the example image:
|
||||
- The existing Anima real workflow rerun retained exact before/after latent and
|
||||
pixel equality.
|
||||
|
||||
Frontend and API validation for a missing source mask names `inpaint_mask`.
|
||||
Supplying an incompatible model produces this exact error:
|
||||
`The provided ControlNet does not use an inpaint source mask; use Apply Advanced
|
||||
ControlNet instead.`
|
||||
Frontend and API validation confirms that the normal Apply node exposes the
|
||||
optional source mask as `inpaint_mask`. When connected to a ControlNet without
|
||||
source-mask support, that input is ignored and the normal control path is used.
|
||||
|
||||
Workflow and result screenshots are linked from the PR instead of stored here
|
||||
to avoid repository growth.
|
||||
|
||||
File diff suppressed because one or more lines are too long
@@ -2,7 +2,7 @@ import os
|
||||
import sys
|
||||
import unittest
|
||||
from types import SimpleNamespace
|
||||
from unittest.mock import patch, sentinel
|
||||
from unittest.mock import Mock, patch, sentinel
|
||||
|
||||
comfyui_path = os.environ.get("COMFYUI_PATH")
|
||||
if comfyui_path:
|
||||
@@ -13,7 +13,7 @@ import torch
|
||||
from comfy.controlnet import T2IAdapter
|
||||
|
||||
from adv_control.control import ControlNetAdvanced, T2IAdapterAdvanced
|
||||
from adv_control.nodes_main import AdvancedControlNetApply, AdvancedControlNetInpaintingApply
|
||||
from adv_control.nodes_main import AdvancedControlNetApply
|
||||
from adv_control.utils import ControlWeights
|
||||
|
||||
|
||||
@@ -157,6 +157,38 @@ class T2IAdapterTests(unittest.TestCase):
|
||||
|
||||
|
||||
class AdvancedControlNetApplyTests(unittest.TestCase):
|
||||
def apply_control(self, concat_mask, image, inpaint_mask, effect_mask=None):
|
||||
control_net = SimpleNamespace(concat_mask=concat_mask, copy=Mock(return_value=sentinel.control_copy))
|
||||
applied_control = SimpleNamespace(
|
||||
allow_condhint_latents=False,
|
||||
require_vae=False,
|
||||
postpone_condhint_latents_check=False,
|
||||
disarm=Mock(),
|
||||
set_cond_hint=Mock(),
|
||||
set_cond_hint_mask=Mock(),
|
||||
set_previous_controlnet=Mock(),
|
||||
verify_all_weights=Mock(),
|
||||
)
|
||||
applied_control.set_cond_hint.return_value = applied_control
|
||||
positive = [[sentinel.positive_tensor, {}]]
|
||||
|
||||
with patch("adv_control.nodes_main.convert_to_advanced", return_value=applied_control), \
|
||||
patch("adv_control.nodes_main.is_advanced_controlnet", return_value=True):
|
||||
AdvancedControlNetApply.execute(
|
||||
positive=positive,
|
||||
negative=[],
|
||||
control_net=control_net,
|
||||
image=image,
|
||||
strength=1.0,
|
||||
start_percent=0.0,
|
||||
end_percent=1.0,
|
||||
mask_optional=effect_mask,
|
||||
vae_optional=sentinel.vae,
|
||||
inpaint_mask=inpaint_mask,
|
||||
)
|
||||
|
||||
return applied_control
|
||||
|
||||
def test_all_zero_effect_mask_returns_original_conditioning(self):
|
||||
positive = [[sentinel.positive_tensor, {"name": "positive"}]]
|
||||
negative = [[sentinel.negative_tensor, {"name": "negative"}]]
|
||||
@@ -175,50 +207,29 @@ class AdvancedControlNetApplyTests(unittest.TestCase):
|
||||
self.assertIs(result.args[0], positive)
|
||||
self.assertIs(result.args[1], negative)
|
||||
|
||||
|
||||
class AdvancedInpaintingApplyTests(unittest.TestCase):
|
||||
def test_source_mask_and_effect_mask_stay_independent(self):
|
||||
image = torch.ones((1, 2, 2, 3))
|
||||
inpaint_mask = torch.tensor([[[1.0, 0.0], [1.0, 0.0]]])
|
||||
effect_mask = torch.full((1, 2, 2), 0.25)
|
||||
control_net = SimpleNamespace(concat_mask=True)
|
||||
applied_control = self.apply_control(True, image, inpaint_mask, effect_mask)
|
||||
|
||||
with patch.object(AdvancedControlNetApply, "execute", return_value=sentinel.output) as apply:
|
||||
result = AdvancedControlNetInpaintingApply.execute(
|
||||
positive=sentinel.positive,
|
||||
negative=sentinel.negative,
|
||||
control_net=control_net,
|
||||
vae=sentinel.vae,
|
||||
image=image,
|
||||
inpaint_mask=inpaint_mask,
|
||||
strength=1.0,
|
||||
start_percent=0.0,
|
||||
end_percent=1.0,
|
||||
effect_mask_optional=effect_mask,
|
||||
)
|
||||
inputs = applied_control.set_cond_hint.call_args.args
|
||||
source_mask = 1.0 - inpaint_mask.unsqueeze(1)
|
||||
torch.testing.assert_close(inputs[0], (image * source_mask.movedim(1, -1)).movedim(-1, 1))
|
||||
torch.testing.assert_close(inputs[4][0], source_mask)
|
||||
torch.testing.assert_close(applied_control.set_cond_hint_mask.call_args.args[0], effect_mask)
|
||||
|
||||
self.assertIs(result, sentinel.output)
|
||||
inputs = apply.call_args.kwargs
|
||||
torch.testing.assert_close(inputs["mask_optional"], effect_mask)
|
||||
torch.testing.assert_close(inputs["extra_concat"][0], 1.0 - inpaint_mask.unsqueeze(1))
|
||||
def test_inpaint_mask_is_ignored_for_other_controlnets(self):
|
||||
image = torch.ones((1, 2, 2, 3))
|
||||
inpaint_mask = torch.tensor([[[1.0, 0.0], [1.0, 0.0]]])
|
||||
applied_control = self.apply_control(False, image, inpaint_mask)
|
||||
|
||||
inputs = applied_control.set_cond_hint.call_args.args
|
||||
torch.testing.assert_close(
|
||||
inputs["image"],
|
||||
torch.tensor([[[[0.0, 0.0, 0.0], [1.0, 1.0, 1.0]], [[0.0, 0.0, 0.0], [1.0, 1.0, 1.0]]]]),
|
||||
inputs[0],
|
||||
image.movedim(-1, 1),
|
||||
)
|
||||
|
||||
def test_non_inpaint_control_has_readable_error(self):
|
||||
with self.assertRaisesRegex(ValueError, "does not use an inpaint source mask"):
|
||||
AdvancedControlNetInpaintingApply.execute(
|
||||
positive=[],
|
||||
negative=[],
|
||||
control_net=SimpleNamespace(concat_mask=False),
|
||||
vae=sentinel.vae,
|
||||
image=torch.ones((1, 2, 2, 3)),
|
||||
inpaint_mask=torch.zeros((1, 2, 2)),
|
||||
strength=1.0,
|
||||
start_percent=0.0,
|
||||
end_percent=1.0,
|
||||
)
|
||||
self.assertEqual(inputs[4], [])
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
|
||||
Reference in New Issue
Block a user