(
self,
hidden_states: torch.FloatTensor,
controlnet_cond: List[torch.tensor],
controlnet_mode: List[torch.tensor],
conditioning_scale: List[float],
encoder_hidden_states: torch.Tensor = None,
pooled_projections: torch.Tensor = None,
timestep: torch.LongTensor = None,
img_ids: torch.Tensor = None,
txt_ids: torch.Tensor = None,
guidance: torch.Tensor = None,
joint_attention_kwargs: Optional[Dict[str, Any]] = None,
return_dict: bool = True,
)
| 432 | self.union = union |
| 433 | |
| 434 | def forward( |
| 435 | self, |
| 436 | hidden_states: torch.FloatTensor, |
| 437 | controlnet_cond: List[torch.tensor], |
| 438 | controlnet_mode: List[torch.tensor], |
| 439 | conditioning_scale: List[float], |
| 440 | encoder_hidden_states: torch.Tensor = None, |
| 441 | pooled_projections: torch.Tensor = None, |
| 442 | timestep: torch.LongTensor = None, |
| 443 | img_ids: torch.Tensor = None, |
| 444 | txt_ids: torch.Tensor = None, |
| 445 | guidance: torch.Tensor = None, |
| 446 | joint_attention_kwargs: Optional[Dict[str, Any]] = None, |
| 447 | return_dict: bool = True, |
| 448 | ) -> Union[FluxControlNetOutput, Tuple]: |
| 449 | # ControlNet-Union with multiple conditions |
| 450 | # only load one ControlNet for saving memories |
| 451 | if len(self.nets) == 1: |
| 452 | controlnet = self.nets[0] |
| 453 | for i, (image, mode, scale) in enumerate(zip(controlnet_cond, controlnet_mode, conditioning_scale)): |
| 454 | |
| 455 | block_samples, single_block_samples = controlnet( |
| 456 | hidden_states=hidden_states, |
| 457 | controlnet_cond=image, |
| 458 | controlnet_mode=mode[:, None], |
| 459 | conditioning_scale=scale, |
| 460 | timestep=timestep, |
| 461 | guidance=guidance, |
| 462 | pooled_projections=pooled_projections, |
| 463 | encoder_hidden_states=encoder_hidden_states, |
| 464 | txt_ids=txt_ids, |
| 465 | img_ids=img_ids, |
| 466 | joint_attention_kwargs=joint_attention_kwargs, |
| 467 | return_dict=return_dict, |
| 468 | ) |
| 469 | |
| 470 | # merge samples |
| 471 | if i == 0: |
| 472 | control_block_samples = block_samples |
| 473 | control_single_block_samples = single_block_samples |
| 474 | else: |
| 475 | if block_samples is not None: |
| 476 | control_block_samples = [ |
| 477 | control_block_sample + block_sample |
| 478 | for control_block_sample, block_sample in zip(control_block_samples, block_samples) |
| 479 | ] |
| 480 | |
| 481 | if single_block_samples is not None: |
| 482 | control_single_block_samples = [ |
| 483 | control_single_block_sample + block_sample |
| 484 | for control_single_block_sample, block_sample in zip( |
| 485 | control_single_block_samples, single_block_samples |
| 486 | ) |
| 487 | ] |
| 488 | |
| 489 | # Regular Multi-ControlNets |
| 490 | # load all ControlNets into memories |
| 491 | else: |
nothing calls this directly
no outgoing calls
no test coverage detected