Merge pull request #257 from Kosinkadink/fix/advanced-inpaint-control

Add modern ControlNet inpainting support
This commit is contained in:
Jedrzej Kosinski
2026-07-28 04:35:32 -07:00
committed by GitHub
9 changed files with 368 additions and 34 deletions
+9 -6
View File
@@ -64,22 +64,22 @@ class ControlNetAdvanced(ControlNet, AdvancedControlBase):
# make cond_hint appropriate dimensions
# TODO: change this to not require cond_hint upscaling every step when self.sub_idxs are present
if self.sub_idxs is not None or self.cond_hint is None or x_noisy.shape[2] * self.real_compression_ratio != self.cond_hint.shape[2] or x_noisy.shape[3] * self.real_compression_ratio != self.cond_hint.shape[3]:
if self.sub_idxs is not None or self.cond_hint is None or x_noisy.shape[-2] * self.real_compression_ratio != self.cond_hint.shape[-2] or x_noisy.shape[-1] * self.real_compression_ratio != self.cond_hint.shape[-1]:
if self.cond_hint is not None:
del self.cond_hint
self.cond_hint = None
self.real_compression_ratio = self.compression_ratio
compression_ratio = self.compression_ratio
if self.vae is not None and self.mult_by_ratio_when_vae:
compression_ratio *= self.vae.downscale_ratio
compression_ratio *= self.vae.spacial_compression_encode()
# if self.cond_hint_original length greater or equal to real latent count, subdivide it before scaling
if self.sub_idxs is not None:
actual_cond_hint_orig = self.cond_hint_original
if self.cond_hint_original.size(0) < self.full_latent_length:
actual_cond_hint_orig = extend_to_batch_size(tensor=actual_cond_hint_orig, batch_size=self.full_latent_length)
self.cond_hint = comfy.utils.common_upscale(actual_cond_hint_orig[self.sub_idxs], x_noisy.shape[3] * compression_ratio, x_noisy.shape[2] * compression_ratio, self.upscale_algorithm, "center")
self.cond_hint = comfy.utils.common_upscale(actual_cond_hint_orig[self.sub_idxs], x_noisy.shape[-1] * compression_ratio, x_noisy.shape[-2] * compression_ratio, self.upscale_algorithm, "center")
else:
self.cond_hint = comfy.utils.common_upscale(self.cond_hint_original, x_noisy.shape[3] * compression_ratio, x_noisy.shape[2] * compression_ratio, self.upscale_algorithm, "center")
self.cond_hint = comfy.utils.common_upscale(self.cond_hint_original, x_noisy.shape[-1] * compression_ratio, x_noisy.shape[-2] * compression_ratio, self.upscale_algorithm, "center")
self.cond_hint = self.preprocess_image(self.cond_hint)
if self.vae is not None:
loaded_models = comfy.model_management.loaded_models(only_currently_used=True)
@@ -93,7 +93,10 @@ class ControlNetAdvanced(ControlNet, AdvancedControlBase):
to_concat = []
for c in self.extra_concat_orig:
c = c.to(self.cond_hint.device)
c = comfy.utils.common_upscale(c, self.cond_hint.shape[3], self.cond_hint.shape[2], self.upscale_algorithm, "center")
c = comfy.utils.common_upscale(c, self.cond_hint.shape[-1], self.cond_hint.shape[-2], self.upscale_algorithm, "center")
if c.ndim < self.cond_hint.ndim:
c = c.unsqueeze(2)
c = comfy.utils.repeat_to_batch_size(c, self.cond_hint.shape[2], dim=2)
to_concat.append(comfy.utils.repeat_to_batch_size(c, self.cond_hint.shape[0]))
self.cond_hint = torch.cat([self.cond_hint] + to_concat, dim=1)
@@ -123,7 +126,7 @@ class ControlNetAdvanced(ControlNet, AdvancedControlBase):
return super().pre_run_advanced(*args, **kwargs)
def apply_advanced_strengths_and_masks(self, x: Tensor, batched_number: int, flux_shape=None):
if self.is_flux:
if self.is_flux or x.ndim == 3:
flux_shape = self.x_noisy_shape
return super().apply_advanced_strengths_and_masks(x, batched_number, flux_shape)
+8 -8
View File
@@ -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)
+4 -4
View File
@@ -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=[
+24 -13
View File
@@ -2,6 +2,7 @@ from comfy_api.latest import io
from torch import Tensor
import folder_paths
import comfy.utils
from .control import load_controlnet, convert_to_advanced, is_advanced_controlnet, is_sd3_advanced_controlnet
from .control_lllite import load_anima_lllite
@@ -16,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)
@@ -42,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)
@@ -96,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),
@@ -112,10 +114,18 @@ 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):
if strength == 0:
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)
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 = {}
@@ -134,7 +144,7 @@ class AdvancedControlNetApply(io.ComfyNode):
if control_net is None:
raise Exception("Passed in control_net is None; something must have went wrong when loading it from a Load ControlNet node.")
# copy, convert to advanced if needed, and set cond
c_net = convert_to_advanced(control_net.copy()).set_cond_hint(control_hint, strength, (start_percent, end_percent), vae_optional)
c_net = convert_to_advanced(control_net.copy()).set_cond_hint(control_hint, strength, (start_percent, end_percent), vae_optional, extra_concat)
if is_advanced_controlnet(c_net):
# disarm node check
c_net.disarm()
@@ -153,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)
@@ -197,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),
@@ -213,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)
+2 -2
View File
@@ -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)
+1 -1
View File
@@ -359,7 +359,7 @@ def prepare_mask_batch(mask: Tensor, shape: Tensor, multiplier: int=1, match_dim
mask = mask.clone()
if flux_shape is not None:
multiplier = multiplier * 0.5
mask = torch.nn.functional.interpolate(mask.reshape((-1, 1, mask.shape[-2], mask.shape[-1])), size=(round(flux_shape[-2]*multiplier), round(flux_shape[-1]*multiplier)), mode="bilinear")
mask = torch.nn.functional.interpolate(mask.reshape((-1, 1, mask.shape[-2], mask.shape[-1])), size=(math.ceil(flux_shape[-2]*multiplier), math.ceil(flux_shape[-1]*multiplier)), mode="bilinear")
mask = rearrange(mask, "b c h w -> b (h w) c")
else:
mask = torch.nn.functional.interpolate(mask.reshape((-1, 1, mask.shape[-2], mask.shape[-1])), size=(round(shape[-2]*multiplier), round(shape[-1]*multiplier)), mode="bilinear")
+83
View File
@@ -0,0 +1,83 @@
# Qwen Image ControlNet inpainting
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**, 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
Download the official inputs to `ComfyUI/input` with these exact names:
- [`acn_qwen_inpaint_source.png`](https://huggingface.co/InstantX/Qwen-Image-ControlNet-Inpainting/resolve/main/assets/images/image1.png)
- [`acn_qwen_inpaint_mask.png`](https://huggingface.co/InstantX/Qwen-Image-ControlNet-Inpainting/resolve/main/assets/masks/mask1.png)
The model author's repository is
[`InstantX/Qwen-Image-ControlNet-Inpainting`](https://huggingface.co/InstantX/Qwen-Image-ControlNet-Inpainting).
Download every model below to the listed folder under `ComfyUI/models`:
| File and exact download | Folder |
| --- | --- |
| [`qwen_image_fp8_e4m3fn.safetensors`](https://huggingface.co/Comfy-Org/Qwen-Image_ComfyUI/resolve/main/split_files/diffusion_models/qwen_image_fp8_e4m3fn.safetensors) | `diffusion_models` |
| [`qwen_2.5_vl_7b_fp8_scaled.safetensors`](https://huggingface.co/Comfy-Org/Qwen-Image_ComfyUI/resolve/main/split_files/text_encoders/qwen_2.5_vl_7b_fp8_scaled.safetensors) | `text_encoders` |
| [`qwen_image_vae.safetensors`](https://huggingface.co/Comfy-Org/Qwen-Image_ComfyUI/resolve/main/split_files/vae/qwen_image_vae.safetensors) | `vae` |
| [`Qwen-Image-InstantX-ControlNet-Inpainting.safetensors`](https://huggingface.co/Comfy-Org/Qwen-Image-InstantX-ControlNets/resolve/main/split_files/controlnet/Qwen-Image-InstantX-ControlNet-Inpainting.safetensors) | `controlnet` |
| [`Qwen-Image-Lightning-4steps-V1.0.safetensors`](https://huggingface.co/lightx2v/Qwen-Image-Lightning/resolve/main/Qwen-Image-Lightning-4steps-V1.0.safetensors) | `loras` (optional and bypassed) |
## Run
1. Download the two inputs and five model files to the folders above.
2. Load `qwen_image_inpainting.json` in ComfyUI.
3. Queue the workflow unchanged.
For command-line input reproduction:
```sh
curl -L https://huggingface.co/InstantX/Qwen-Image-ControlNet-Inpainting/resolve/main/assets/images/image1.png -o ComfyUI/input/acn_qwen_inpaint_source.png
curl -L https://huggingface.co/InstantX/Qwen-Image-ControlNet-Inpainting/resolve/main/assets/masks/mask1.png -o ComfyUI/input/acn_qwen_inpaint_mask.png
```
The unchanged example uses seed `134554158057228` (fixed), 20 steps, CFG 2.5,
Euler, the simple scheduler, denoise 1.0, model shift 3.1, control strength 1.0,
and control start/end 0.0/1.0. Its prompt is `The Queen, on a throne,
surrounded by Knights, HD, Realistic, Octane Render, Unreal engine`; the
negative prompt is one space. The source is scaled with area interpolation to a
maximum dimension of 1536. The optional 4-step LoRA remains bypassed; enabling
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` 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
These results were measured with fixed inputs and settings; they are recorded
here rather than inferred from the example image:
- A fresh isolated vanilla-versus-Advanced run had latent maximum/mean absolute
differences `0/0`, pixel maximum/mean differences `0/0`, and 0 changed
pixels.
- An all-one effect mask exactly equaled unmasked Advanced output at latent and
pixel level. An all-zero effect mask exactly equaled no ControlNet at latent
and pixel level.
- For right-half token-mask injection, relative to full control the left latent
mean delta was `0` and the right was `0.1332103`; relative to no control the
left was `0` and the right was `0.1835042`.
- In a per-latent batch, the sample with strength 0 exactly equaled no control
at latent and pixel level.
- Soft weights, timestep scheduling, and two-control stacking each executed
successfully.
- The existing Anima real workflow rerun retained exact before/after latent and
pixel equality.
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
+236
View File
@@ -0,0 +1,236 @@
import os
import sys
import unittest
from types import SimpleNamespace
from unittest.mock import Mock, patch, sentinel
comfyui_path = os.environ.get("COMFYUI_PATH")
if comfyui_path:
sys.path.insert(0, comfyui_path)
import torch
from comfy.controlnet import T2IAdapter
from adv_control.control import ControlNetAdvanced, T2IAdapterAdvanced
from adv_control.nodes_main import AdvancedControlNetApply
from adv_control.utils import ControlWeights
class StopControlModel(Exception):
pass
class ControlModel:
dtype = torch.float32
def __init__(self):
self.hint = None
def __call__(self, x, hint, timesteps, context, **kwargs):
self.hint = hint
raise StopControlModel
class VideoVAE:
downscale_ratio = (4, 8, 8)
def __init__(self):
self.encoded_shape = None
def spacial_compression_encode(self):
return 8
def encode(self, image):
self.encoded_shape = image.shape
return torch.ones((image.shape[0], 4, 2, 2, 2))
class ModernControlPreprocessingTests(unittest.TestCase):
def test_effect_mask_is_resized_to_qwen_tokens(self):
control = ControlNetAdvanced(ControlModel(), None)
control.x_noisy_shape = (1, 16, 4, 6)
control.mask_cond_hint = torch.tensor(
[[[[0.0, 0.0, 0.0, 1.0, 1.0, 1.0]] * 4]]
)
control.tk_mask_cond_hint = None
control.weights = SimpleNamespace(has_uncond_multiplier=False, has_uncond_mask=False)
control.latent_keyframes = None
control._current_timestep_keyframe = SimpleNamespace(strength=1.0)
output = torch.ones((1, 6, 4))
control.apply_advanced_strengths_and_masks(output, batched_number=1)
expected = torch.tensor(
[[[0.0] * 4, [0.5] * 4, [1.0] * 4, [0.0] * 4, [0.5] * 4, [1.0] * 4]]
)
torch.testing.assert_close(output, expected)
def test_effect_mask_matches_padded_flux_tokens_for_odd_latent_size(self):
control = ControlNetAdvanced(ControlModel(), None)
control.x_noisy_shape = (1, 16, 5, 7)
control.mask_cond_hint = torch.ones((1, 1, 5, 7))
control.tk_mask_cond_hint = None
control.weights = SimpleNamespace(has_uncond_multiplier=False, has_uncond_mask=False)
control.latent_keyframes = None
control._current_timestep_keyframe = SimpleNamespace(strength=1.0)
output = torch.ones((1, 12, 4))
control.apply_advanced_strengths_and_masks(output, batched_number=1)
torch.testing.assert_close(output, torch.ones_like(output))
def test_vae_compression_and_source_mask_match_5d_hint(self):
control_model = ControlModel()
vae = VideoVAE()
control = ControlNetAdvanced(control_model, None, compression_ratio=1, latent_format=SimpleNamespace(process_in=lambda value: value))
control.real_compression_ratio = 1
control.cond_hint_original = torch.ones((1, 3, 16, 16))
control.cond_hint = None
control.vae = vae
control.extra_concat_orig = [torch.zeros((1, 1, 16, 16))]
control.sub_idxs = None
control.model_sampling_current = SimpleNamespace(timestep=lambda value: value, calculate_input=lambda timestep, value: value)
control.prepare_mask_cond_hint = lambda **kwargs: None
with self.assertRaises(StopControlModel):
control.sliding_get_control(
torch.ones((1, 4, 2, 2, 2)),
torch.ones(1),
{"c_crossattn": torch.ones((1, 1, 1))},
1,
{},
)
self.assertEqual(tuple(vae.encoded_shape), (1, 16, 16, 3))
self.assertEqual(tuple(control_model.hint.shape), (1, 5, 2, 2, 2))
class T2IAdapterTests(unittest.TestCase):
def test_effect_masks_are_applied_to_adapter_features(self):
control = T2IAdapterAdvanced(SimpleNamespace(), None, channels_in=3)
control.weights = ControlWeights.t2iadapter()
control.latent_keyframes = None
control.tk_mask_cond_hint = None
control._current_timestep_keyframe = SimpleNamespace(strength=1.0)
masks = {
"zero": torch.zeros((1, 1, 8, 8)),
"one": torch.ones((1, 1, 8, 8)),
"half": torch.cat((torch.zeros((1, 1, 8, 4)), torch.ones((1, 1, 8, 4))), dim=3),
}
for name, mask in masks.items():
with self.subTest(name=name):
features = torch.ones((1, 4, 8, 8))
control.mask_cond_hint = mask
control.apply_advanced_strengths_and_masks(features, batched_number=1)
torch.testing.assert_close(features, mask.expand_as(features))
def test_sliding_context_extends_single_hint_to_full_latent_length(self):
control = T2IAdapterAdvanced(SimpleNamespace(), None, channels_in=3)
original_hint = torch.ones((1, 3, 8, 8))
control.cond_hint_original = original_hint
control.cond_hint = None
control.sub_idxs = [2, 3]
control.full_latent_length = 4
control.prepare_mask_cond_hint = lambda **kwargs: None
selected_hint = None
def get_control(adapter, *args, **kwargs):
nonlocal selected_hint
selected_hint = adapter.cond_hint_original.clone()
return sentinel.output
with patch.object(T2IAdapter, "get_control", get_control):
result = control.get_control_advanced(
torch.ones((2, 4, 8, 8)),
torch.ones(2),
{},
1,
{},
)
self.assertIs(result, sentinel.output)
self.assertEqual(tuple(selected_hint.shape), (2, 3, 8, 8))
torch.testing.assert_close(selected_hint, original_hint.repeat(2, 1, 1, 1))
self.assertIs(control.cond_hint_original, original_hint)
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"}]]
result = AdvancedControlNetApply.execute(
positive=positive,
negative=negative,
control_net=sentinel.control_net,
image=torch.ones((1, 8, 8, 3)),
strength=1.0,
start_percent=0.0,
end_percent=1.0,
mask_optional=torch.zeros((1, 8, 8)),
)
self.assertIs(result.args[0], positive)
self.assertIs(result.args[1], negative)
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)
applied_control = self.apply_control(True, image, inpaint_mask, 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)
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[0],
image.movedim(-1, 1),
)
self.assertEqual(inputs[4], [])
if __name__ == "__main__":
unittest.main()