(self, hidden_states, context=None, mask=None)
| 68 | nn.Conv2d(inner_dim, query_dim, kernel_size=1, bias=True)) |
| 69 | |
| 70 | def forward(self, hidden_states, context=None, mask=None): |
| 71 | if self.training: |
| 72 | raise NotImplementedError(WARN_MSG) |
| 73 | |
| 74 | batch_size, dim, _, sequence_length = hidden_states.shape |
| 75 | |
| 76 | q = self.to_q(hidden_states) |
| 77 | context = context if context is not None else hidden_states |
| 78 | k = self.to_k(context) |
| 79 | v = self.to_v(context) |
| 80 | |
| 81 | # Validate mask |
| 82 | if mask is not None: |
| 83 | expected_mask_shape = [batch_size, sequence_length, 1, 1] |
| 84 | if mask.dtype == torch.bool: |
| 85 | mask = mask.logical_not().float() * -1e4 |
| 86 | elif mask.dtype == torch.int64: |
| 87 | mask = (1 - mask).float() * -1e4 |
| 88 | elif mask.dtype != torch.float32: |
| 89 | raise TypeError(f"Unexpected dtype for mask: {mask.dtype}") |
| 90 | |
| 91 | if len(mask.size()) == 2: |
| 92 | mask = mask.unsqueeze(2).unsqueeze(2) |
| 93 | |
| 94 | if list(mask.size()) != expected_mask_shape: |
| 95 | raise RuntimeError( |
| 96 | f"Invalid shape for `mask` (Expected {expected_mask_shape}, got {list(mask.size())}" |
| 97 | ) |
| 98 | |
| 99 | if ATTENTION_IMPLEMENTATION_IN_EFFECT == AttentionImplementations.ORIGINAL: |
| 100 | attn = attention.original(q, k, v, mask, self.heads, self.dim_head) |
| 101 | |
| 102 | elif ATTENTION_IMPLEMENTATION_IN_EFFECT == AttentionImplementations.SPLIT_EINSUM: |
| 103 | attn = attention.split_einsum(q, k, v, mask, self.heads, self.dim_head) |
| 104 | |
| 105 | elif ATTENTION_IMPLEMENTATION_IN_EFFECT == AttentionImplementations.SPLIT_EINSUM_V2: |
| 106 | attn = attention.split_einsum_v2(q, k, v, mask, self.heads, self.dim_head) |
| 107 | |
| 108 | else: |
| 109 | raise ValueError(ATTENTION_IMPLEMENTATION_IN_EFFECT) |
| 110 | |
| 111 | return self.to_out(attn) |
| 112 | |
| 113 | |
| 114 | def linear_to_conv2d_map(state_dict, prefix, local_metadata, strict, |
nothing calls this directly
no outgoing calls
no test coverage detected