(self, query, key, value, attn_mask = None)
| 186 | self.norm_factor = math.sqrt(self.hidden_size_per_attention_head) |
| 187 | |
| 188 | def forward(self, query, key, value, attn_mask = None): |
| 189 | # query/key/value: [sq, b, h] |
| 190 | sq, b, _ = query.size() |
| 191 | |
| 192 | assert torch.allclose(query, key), 'Only Support Self-Attention Currently' |
| 193 | sk = sq |
| 194 | mixed_x_layer = self.in_proj(query) |
| 195 | |
| 196 | # [sq, b, (np * 3 * hn)] --> [sq, b, np, 3 * hn] |
| 197 | new_tensor_shape = mixed_x_layer.size()[:-1] + \ |
| 198 | (self.num_attention_heads_per_partition, |
| 199 | 3 * self.hidden_size_per_attention_head) |
| 200 | mixed_x_layer = mixed_x_layer.view(*new_tensor_shape) |
| 201 | |
| 202 | # [sq, b, np, 3 * hn] --> 3 [sq, b, np, hn] |
| 203 | query_layer, key_layer, value_layer = mixed_x_layer.split( |
| 204 | self.hidden_size_per_attention_head, dim=-1) |
| 205 | |
| 206 | # [sq, b, np, hn] -> [sq, b * np, hn] |
| 207 | query_layer = query_layer.view(sq, |
| 208 | b * self.num_attention_heads_per_partition, |
| 209 | self.hidden_size_per_attention_head).transpose(0, 1) |
| 210 | # [sk, b, np, hn] -> [sk, b * np, hn] |
| 211 | key_layer = key_layer.view(sk, |
| 212 | b * self.num_attention_heads_per_partition, |
| 213 | self.hidden_size_per_attention_head).transpose(0, 1) |
| 214 | |
| 215 | q_scaled = query_layer / self.norm_factor |
| 216 | if attn_mask is not None: |
| 217 | attention_probs = torch.baddbmm(attn_mask, q_scaled, key_layer.transpose(-2, -1)) |
| 218 | else: |
| 219 | attention_probs = torch.bmm(q_scaled, key_layer.transpose(-2, -1)) |
| 220 | attention_probs = attention_probs.softmax(dim=-1) |
| 221 | |
| 222 | value_layer = value_layer.view(sk, |
| 223 | b * self.num_attention_heads_per_partition, |
| 224 | self.hidden_size_per_attention_head).transpose(0, 1) |
| 225 | |
| 226 | # matmul: [b * np, sq, hn] |
| 227 | context_layer = torch.bmm(attention_probs, value_layer) |
| 228 | |
| 229 | # change view [b, np, sq, hn] |
| 230 | context_layer = context_layer.view(b, |
| 231 | self.num_attention_heads_per_partition, |
| 232 | sq, self.hidden_size_per_attention_head) |
| 233 | |
| 234 | # [b, np, sq, hn] --> [sq, b, np, hn] |
| 235 | context_layer = context_layer.permute(2, 0, 1, 3).contiguous() |
| 236 | |
| 237 | # [sq, b, np, hn] --> [sq, b, hp] |
| 238 | new_context_layer_shape = context_layer.size()[:-2] + \ |
| 239 | (self.hidden_size_per_partition,) |
| 240 | context_layer = context_layer.view(*new_context_layer_shape) |
| 241 | |
| 242 | output = self.out_proj(context_layer) |
| 243 | |
| 244 | return output |
| 245 |
nothing calls this directly
no outgoing calls
no test coverage detected