(query_layer, key_layer, value_layer, attention_mask,
attention_dropout=None, log_attention_weights=None, scaling_attention_score=True, **kwargs)
| 17 | import contextlib |
| 18 | |
| 19 | def standard_attention(query_layer, key_layer, value_layer, attention_mask, |
| 20 | attention_dropout=None, log_attention_weights=None, scaling_attention_score=True, **kwargs): |
| 21 | # We disable the PB-relax-Attention and only changes the order of computation, because it is enough for most of training. |
| 22 | # The implementation in the paper can be done very easily, if you really need it to train very deep transformers. |
| 23 | |
| 24 | if scaling_attention_score: |
| 25 | query_layer = query_layer / math.sqrt(query_layer.shape[-1]) |
| 26 | attention_scores = torch.matmul(query_layer, key_layer.transpose(-1, -2)) |
| 27 | if log_attention_weights is not None: |
| 28 | attention_scores += log_attention_weights |
| 29 | |
| 30 | if not (attention_mask.shape[-2] == 1 and (attention_mask > 0).all()): |
| 31 | # if auto-regressive, skip |
| 32 | attention_scores = torch.mul(attention_scores, attention_mask) - \ |
| 33 | 10000.0 * (1.0 - attention_mask) |
| 34 | |
| 35 | attention_probs = F.softmax(attention_scores, dim=-1) |
| 36 | |
| 37 | if attention_dropout is not None: |
| 38 | if mpu.get_cuda_rng_tracker is not None: |
| 39 | with mpu.get_cuda_rng_tracker().fork(): |
| 40 | attention_probs = attention_dropout(attention_probs) |
| 41 | else: |
| 42 | attention_probs = attention_dropout(attention_probs) |
| 43 | |
| 44 | context_layer = torch.matmul(attention_probs, value_layer) |
| 45 | return context_layer |
| 46 | |
| 47 | def attention_fn_default(query_layer, key_layer, value_layer, attention_mask, |
| 48 | attention_dropout=None, log_attention_weights=None, scaling_attention_score=True, **kwargs): |
no outgoing calls
no test coverage detected