MCPcopy Create free account
hub / github.com/Shakker-Labs/RepText / forward

Method forward

controlnet_flux.py:434–529  ·  view source on GitHub ↗
(
        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,
    )

Source from the content-addressed store, hash-verified

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:

Callers

nothing calls this directly

Calls

no outgoing calls

Tested by

no test coverage detected