MCPcopy Create free account
hub / github.com/OpenDriveLab/ReSim / standard_attention

Function standard_attention

SwissArmyTransformer/sat/transformer_defaults.py:19–45  ·  view source on GitHub ↗
(query_layer, key_layer, value_layer, attention_mask,
                       attention_dropout=None, log_attention_weights=None, scaling_attention_score=True, **kwargs)

Source from the content-addressed store, hash-verified

17import contextlib
18
19def 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
47def attention_fn_default(query_layer, key_layer, value_layer, attention_mask,
48 attention_dropout=None, log_attention_weights=None, scaling_attention_score=True, **kwargs):

Callers 1

attention_fn_defaultFunction · 0.85

Calls

no outgoing calls

Tested by

no test coverage detected