| 144 | |
| 145 | |
| 146 | class 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 = [] |
no outgoing calls
no test coverage detected