Top query self-attention layer abstract class. Self-attention layer takes input with size [b, s, h] and returns output of the same size.
| 277 | |
| 278 | |
| 279 | class TopQuerySelfAttention(torch.nn.Module): |
| 280 | """Top query self-attention layer abstract class. |
| 281 | Self-attention layer takes input with size [b, s, h] |
| 282 | and returns output of the same size. |
| 283 | """ |
| 284 | |
| 285 | def __init__( |
| 286 | self, |
| 287 | hidden_size, |
| 288 | num_attention_heads, |
| 289 | layer_number, |
| 290 | fp16=True, |
| 291 | attention_softmax_in_fp32=True, |
| 292 | ): |
| 293 | super(TopQuerySelfAttention, self).__init__() |
| 294 | self.hidden_size = hidden_size |
| 295 | self.num_attention_heads = num_attention_heads |
| 296 | self.fp16 = fp16 |
| 297 | self.attention_softmax_in_fp32 = attention_softmax_in_fp32 |
| 298 | self.layer_number = max(1, layer_number) |
| 299 | |
| 300 | assert self.hidden_size % self.num_attention_heads == 0 |
| 301 | self.hidden_size_per_attention_head = int(self.hidden_size // self.num_attention_heads) |
| 302 | |
| 303 | self.query = torch.nn.Linear(self.hidden_size, self.hidden_size) |
| 304 | self.key = torch.nn.Linear(self.hidden_size, self.hidden_size) |
| 305 | self.value = torch.nn.Linear(self.hidden_size, self.hidden_size) |
| 306 | |
| 307 | self.norm_factor = math.sqrt(self.hidden_size_per_attention_head) |
| 308 | self.softmax = torch.nn.Softmax(dim=-1) |
| 309 | |
| 310 | self.dense = torch.nn.Linear(self.hidden_size, self.hidden_size) |
| 311 | |
| 312 | def forward( |
| 313 | self, |
| 314 | hidden_states, |
| 315 | query_hidden_state, |
| 316 | attention_mask, |
| 317 | layer_past=None, |
| 318 | get_key_value=False, |
| 319 | prompt_length=None, |
| 320 | context_length=None, |
| 321 | ): |
| 322 | |
| 323 | # hidden_states: [sq, b, h] |
| 324 | if hasattr(torch._C, 'grouped_matmul_bias') and not isinstance(self.query, QuantizedLinear): |
| 325 | query_layer, key_layer, value_layer = torch._C.grouped_matmul_bias([query_hidden_state, hidden_states, hidden_states], |
| 326 | [self.query.weight, self.key.weight, self.value.weight], |
| 327 | [self.query.bias, self.key.bias, self.value.bias]) |
| 328 | else: |
| 329 | query_layer = self.query(query_hidden_state) |
| 330 | key_layer = self.key(hidden_states) |
| 331 | value_layer = self.value(hidden_states) |
| 332 | |
| 333 | fallback = not hasattr(torch._C, 'fused_multi_head_attention_inference_v2') |
| 334 | |
| 335 | if fallback: |
| 336 | if hasattr(torch._C, 'fused_codegeex_qkv_reshape'): |