(self, hidden_states, encoder_hidden_states=None, attn_mask=None, ipadapter_kwargs=None, qkv_preprocessor=None)
| 35 | return hidden_states |
| 36 | |
| 37 | def torch_forward(self, hidden_states, encoder_hidden_states=None, attn_mask=None, ipadapter_kwargs=None, qkv_preprocessor=None): |
| 38 | if encoder_hidden_states is None: |
| 39 | encoder_hidden_states = hidden_states |
| 40 | |
| 41 | batch_size = encoder_hidden_states.shape[0] |
| 42 | |
| 43 | q = self.to_q(hidden_states) |
| 44 | k = self.to_k(encoder_hidden_states) |
| 45 | v = self.to_v(encoder_hidden_states) |
| 46 | |
| 47 | q = q.view(batch_size, -1, self.num_heads, self.head_dim).transpose(1, 2) |
| 48 | k = k.view(batch_size, -1, self.num_heads, self.head_dim).transpose(1, 2) |
| 49 | v = v.view(batch_size, -1, self.num_heads, self.head_dim).transpose(1, 2) |
| 50 | |
| 51 | if qkv_preprocessor is not None: |
| 52 | q, k, v = qkv_preprocessor(q, k, v) |
| 53 | |
| 54 | hidden_states = torch.nn.functional.scaled_dot_product_attention(q, k, v, attn_mask=attn_mask) |
| 55 | if ipadapter_kwargs is not None: |
| 56 | hidden_states = self.interact_with_ipadapter(hidden_states, q, **ipadapter_kwargs) |
| 57 | hidden_states = hidden_states.transpose(1, 2).reshape(batch_size, -1, self.num_heads * self.head_dim) |
| 58 | hidden_states = hidden_states.to(q.dtype) |
| 59 | |
| 60 | hidden_states = self.to_out(hidden_states) |
| 61 | |
| 62 | return hidden_states |
| 63 | |
| 64 | def xformers_forward(self, hidden_states, encoder_hidden_states=None, attn_mask=None): |
| 65 | if encoder_hidden_states is None: |
no test coverage detected