Support 'none' type in Load ControlNet++ Model (Single) node to reproduce vanilla ComfyUI results (no type is assigned), more cleanup

This commit is contained in:
Jedrzej Kosinski
2024-07-15 02:42:11 -05:00
parent 33edd4e3ff
commit fa48109769
2 changed files with 25 additions and 71 deletions
+20 -69
View File
@@ -36,8 +36,10 @@ class PlusPlusType:
SEGMENT = "segment"
TILE = "tile"
REPAINT = "inpaint/outpaint"
NONE = "none"
_LIST_WITH_NONE = [OPENPOSE, DEPTH, THICKLINE, THINLINE, NORMAL, SEGMENT, TILE, REPAINT, NONE]
_LIST = [OPENPOSE, DEPTH, THICKLINE, THINLINE, NORMAL, SEGMENT, TILE, REPAINT]
_DICT = {OPENPOSE: 0, DEPTH: 1, THICKLINE: 2, THINLINE: 3, NORMAL: 4, SEGMENT: 5, TILE: 6, REPAINT: 7}
_DICT = {OPENPOSE: 0, DEPTH: 1, THICKLINE: 2, THINLINE: 3, NORMAL: 4, SEGMENT: 5, TILE: 6, REPAINT: 7, NONE: -1}
@classmethod
def to_idx(cls, control_type: str):
@@ -134,6 +136,8 @@ class ControlAddEmbeddingAdv(nn.Module):
self.linear_2 = operations.Linear(out_dim, out_dim, dtype=dtype, device=device)
def forward(self, control_type, dtype, device):
if control_type is None:
control_type = torch.zeros((self.num_control_type,), device=device)
c_type = timestep_embedding(control_type.flatten(), self.in_dim, repeat_only=False).to(dtype).reshape((-1, self.num_control_type * self.in_dim))
return self.linear_2(torch.nn.functional.silu(self.linear_1(c_type)))
@@ -185,13 +189,14 @@ class ControlNetPlusPlus(ControlNetCLDM):
guided_hint = None
if self.control_add_embedding is not None:
control_type = kwargs.get("control_type", [])
control_type = kwargs.get("control_type", None)
emb += self.control_add_embedding(control_type, emb.dtype, emb.device)
guided_hint = self.union_controlnet_merge(hint, control_type, emb, context)
if control_type is not None:
guided_hint = self.union_controlnet_merge(hint, control_type, emb, context)
if guided_hint is None:
guided_hint = self.input_hint_block(hint, emb, context)
guided_hint = self.input_hint_block(hint[0], emb, context)
out_output = []
out_middle = []
@@ -216,68 +221,6 @@ class ControlNetPlusPlus(ControlNetCLDM):
return {"middle": out_middle, "output": out_output}
def forward_bad(self, x: Tensor, hint: Tensor, timesteps, context, control_type: Tensor, y: Tensor=None, **kwargs):
t_emb = timestep_embedding(timesteps, self.model_channels, repeat_only=False).to(x.dtype)
emb = self.time_embed(t_emb)
# inject control type info to time embedding to distinguish different control conditions
control_embeds = timestep_embedding(control_type.flatten(), self.control_add_embed_dim, repeat_only=False).reshape((t_emb.shape[0], -1)).to(t_emb.dtype)
control_emb = self.control_add_embedding(control_embeds)
emb = emb + control_emb
out_output = []
out_middle = []
hs = []
if self.num_classes is not None:
assert y.shape[0] == x.shape[0]
emb = emb + self.label_emb(y)
guided_hint = True
h = x
for module, zero_conv in zip(self.input_blocks, self.zero_convs):
if guided_hint is not None:
h = module(h, emb, context)
# add single/multi conditions to input image
# Condition Transformer provides an easy and effective way to fuse different features naturally
guided_hint = torch.zeros_like(h)
indexes = torch.nonzero(control_type[0])
inputs = []
condition_list = []
for idx in range(indexes.shape[0] + 1):
if idx == indexes.shape[0]:
raw_hint = h
feat_seq = torch.mean(raw_hint, dim=(2, 3)) # N * C
else:
raw_hint = self.input_hint_block(hint[indexes[idx][0]], emb, context)
feat_seq = torch.mean(raw_hint, dim=(2, 3)) # N * C
feat_seq = feat_seq + self.task_embedding[indexes[idx][0]]
feat_seq = feat_seq.repeat((x.shape[0] // feat_seq.shape[0], 1))
inputs.append(feat_seq.unsqueeze(1))
condition_list.append(raw_hint)
z = torch.cat(inputs, dim=1) # NxLxC
z = self.transformer_layes(z)
for idx in range(indexes.shape[0]):
alpha: Tensor = self.spatial_ch_projs(z[:, idx])
alpha = alpha.unsqueeze(-1).unsqueeze(-1)
guided_hint += condition_list[idx] + alpha
# usual controlnet functionality
h += guided_hint
guided_hint = None
else:
h = module(h, emb, context)
out_output.append(zero_conv(h, emb, context))
h = self.middle_block(h, emb, context)
out_middle.append(self.middle_block_out(h, emb, context))
return {"middle": out_middle, "output": out_output}
class ControlNetPlusPlusAdvanced(ControlNet, AdvancedControlBase):
def __init__(self, control_model: ControlNetPlusPlus, timestep_keyframes: TimestepKeyframeGroup, global_average_pooling=False, device=None, load_device=None, manual_cast_dtype=None):
@@ -356,7 +299,11 @@ class ControlNetPlusPlusAdvanced(ControlNet, AdvancedControlBase):
# unlike normal controlnet, need to handle each input image tensor (for each type)
for pp_type, pp_input in self.cond_hint_original.controls.items():
pp_idx = PlusPlusType.to_idx(pp_type)
self.cond_hint_types[pp_idx] = pp_input.strength
# if negative, means no type should be selected (single only)
if pp_idx < 0:
pp_idx = 0
else:
self.cond_hint_types[pp_idx] = pp_input.strength
# if self.cond_hint_original lengths greater or equal to latent count, subdivide
if self.sub_idxs is not None:
actual_cond_hint_orig = pp_input.image
@@ -368,12 +315,16 @@ class ControlNetPlusPlusAdvanced(ControlNet, AdvancedControlBase):
self.cond_hint[pp_idx] = self.cond_hint[pp_idx].to(device=self.device, dtype=dtype)
self.cond_hint_shape = self.cond_hint[pp_idx].shape
# prepare cond_hint_controls to match batchsize
self.cond_hint_types = self.cond_hint_types.unsqueeze(0).to(device=self.device, dtype=dtype).repeat(x_noisy.shape[0], 1)
if self.cond_hint_types.count_nonzero() == 0:
self.cond_hint_types = None
else:
self.cond_hint_types = self.cond_hint_types.unsqueeze(0).to(device=self.device, dtype=dtype).repeat(x_noisy.shape[0], 1)
for i in range(len(self.cond_hint)):
if self.cond_hint[i] is not None:
if x_noisy.shape[0] != self.cond_hint[i].shape[0]:
self.cond_hint[i] = broadcast_image_to_extend(self.cond_hint[i], x_noisy.shape[0], batched_number)
self.cond_hint_types = broadcast_image_to_extend(self.cond_hint_types, x_noisy.shape[0], batched_number)
if self.cond_hint_types is not None:
self.cond_hint_types = broadcast_image_to_extend(self.cond_hint_types, x_noisy.shape[0], batched_number)
# prepare mask_cond_hint
self.prepare_mask_cond_hint(x_noisy=x_noisy, t=t, cond=cond, batched_number=batched_number, dtype=dtype)
+5 -2
View File
@@ -1,4 +1,5 @@
from torch import Tensor
import math
import folder_paths
@@ -34,7 +35,7 @@ class PlusPlusLoaderSingle:
return {
"required": {
"name": (folder_paths.get_filename_list("controlnet"), ),
"control_type": (PlusPlusType._LIST,),
"control_type": (PlusPlusType._LIST_WITH_NONE,),
}
}
@@ -61,7 +62,7 @@ class PlusPlusInputNode:
},
"optional": {
"prev_plus_input": ("PLUS_INPUT",),
"strength": ("FLOAT", {"default": 1.0, "min": 0.0, "max": BIGMAX, "step": 0.01}),
#"strength": ("FLOAT", {"default": 1.0, "min": 0.0, "max": BIGMAX, "step": 0.01}),
}
}
@@ -75,6 +76,8 @@ class PlusPlusInputNode:
prev_plus_input = PlusPlusInputGroup()
prev_plus_input = prev_plus_input.clone()
if math.isclose(strength, 0.0):
strength = 0.0000001
pp_input = PlusPlusInput(image, control_type, strength)
prev_plus_input.add(pp_input)