MCPcopy Create free account
hub / github.com/Kosinkadink/ComfyUI-Advanced-ControlNet / ControlNetPlusPlus

Class ControlNetPlusPlus

adv_control/control_plusplus.py:146–223  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

144
145
146class ControlNetPlusPlus(ControlNetCLDM):
147 def __init__(self, *args,**kwargs):
148 super().__init__(*args, **kwargs)
149
150 operations: comfy.ops.disable_weight_init = kwargs.get("operations", comfy.ops.disable_weight_init)
151 device = kwargs.get("device", None)
152
153 time_embed_dim = self.model_channels * 4
154 control_add_embed_dim = 256
155
156 self.control_add_embedding = ControlAddEmbeddingAdv(control_add_embed_dim, time_embed_dim, self.num_control_type, dtype=self.dtype, device=device, operations=operations)
157
158 def union_controlnet_merge(self, hint: list[Tensor], control_type, emb, context):
159 # Equivalent to: https://github.com/xinsir6/ControlNetPlus/tree/main
160 indexes = torch.nonzero(control_type[0])
161 inputs = []
162 condition_list = []
163
164 for idx in range(indexes.shape[0]):
165 controlnet_cond = self.input_hint_block(hint[indexes[idx][0]], emb, context)
166 feat_seq = torch.mean(controlnet_cond, dim=(2, 3))
167 if idx < indexes.shape[0]:
168 feat_seq += self.task_embedding[indexes[idx][0]].to(dtype=feat_seq.dtype, device=feat_seq.device)
169
170 inputs.append(feat_seq.unsqueeze(1))
171 condition_list.append(controlnet_cond)
172
173 x = torch.cat(inputs, dim=1)
174 x = self.transformer_layes(x)
175
176 controlnet_cond_fuser = None
177 for idx in range(indexes.shape[0]):
178 alpha = self.spatial_ch_projs(x[:, idx])
179 alpha = alpha.unsqueeze(-1).unsqueeze(-1)
180 o = condition_list[idx] + alpha
181 if controlnet_cond_fuser is None:
182 controlnet_cond_fuser = o
183 else:
184 controlnet_cond_fuser += o
185 return controlnet_cond_fuser
186
187 def forward(self, x: Tensor, hint: list[Tensor], timesteps, context, y: Tensor=None, **kwargs):
188 t_emb = timestep_embedding(timesteps, self.model_channels, repeat_only=False).to(x.dtype)
189 emb = self.time_embed(t_emb)
190
191 guided_hint = None
192 if self.control_add_embedding is not None:
193 control_type = kwargs.get("control_type", None)
194
195 emb += self.control_add_embedding(control_type, emb.dtype, emb.device)
196 if control_type is not None:
197 guided_hint = self.union_controlnet_merge(hint, control_type, emb, context)
198
199 if guided_hint is None:
200 guided_hint = self.input_hint_block(hint[0], emb, context)
201
202 out_output = []
203 out_middle = []

Callers 1

load_controlnetplusplusFunction · 0.85

Calls

no outgoing calls

Tested by

no test coverage detected