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,
)
| 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""" |
nothing calls this directly
no outgoing calls
no test coverage detected