MCPcopy Create free account
hub / github.com/UVA-Computer-Vision-Lab/FrameINO / forward

Method forward

architecture/attention_processor.py:556–600  ·  view source on GitHub ↗

r""" The forward method of the `Attention` class. Args: hidden_states (`torch.Tensor`): The hidden states of the query. encoder_hidden_states (`torch.Tensor`, *optional*): The hidden states of the encoder. attention

(
        self,
        hidden_states: torch.Tensor,
        encoder_hidden_states: Optional[torch.Tensor] = None,
        attention_mask: Optional[torch.Tensor] = None,
        **cross_attention_kwargs,
    )

Source from the content-addressed store, hash-verified

554 return self.processor
555
556 def forward(
557 self,
558 hidden_states: torch.Tensor,
559 encoder_hidden_states: Optional[torch.Tensor] = None,
560 attention_mask: Optional[torch.Tensor] = None,
561 **cross_attention_kwargs,
562 ) -> torch.Tensor:
563 r"""
564 The forward method of the `Attention` class.
565
566 Args:
567 hidden_states (`torch.Tensor`):
568 The hidden states of the query.
569 encoder_hidden_states (`torch.Tensor`, *optional*):
570 The hidden states of the encoder.
571 attention_mask (`torch.Tensor`, *optional*):
572 The attention mask to use. If `None`, no mask is applied.
573 **cross_attention_kwargs:
574 Additional keyword arguments to pass along to the cross attention.
575
576 Returns:
577 `torch.Tensor`: The output of the attention layer.
578 """
579 # The `Attention` class can call different attention processors / attention functions
580 # here we simply pass along all tensors to the selected processor class
581 # For standard processors that are defined here, `**cross_attention_kwargs` is empty
582
583 attn_parameters = set(inspect.signature(self.processor.__call__).parameters.keys())
584 quiet_attn_parameters = {"ip_adapter_masks", "ip_hidden_states"}
585 unused_kwargs = [
586 k for k, _ in cross_attention_kwargs.items() if k not in attn_parameters and k not in quiet_attn_parameters
587 ]
588 if len(unused_kwargs) > 0:
589 logger.warning(
590 f"cross_attention_kwargs {unused_kwargs} are not expected by {self.processor.__class__.__name__} and will be ignored."
591 )
592 cross_attention_kwargs = {k: w for k, w in cross_attention_kwargs.items() if k in attn_parameters}
593
594 return self.processor(
595 self,
596 hidden_states,
597 encoder_hidden_states=encoder_hidden_states,
598 attention_mask=attention_mask,
599 **cross_attention_kwargs,
600 )
601
602 def batch_to_head_dim(self, tensor: torch.Tensor) -> torch.Tensor:
603 r"""

Callers

nothing calls this directly

Calls

no outgoing calls

Tested by

no test coverage detected