MCPcopy Create free account
hub / github.com/NJUNLP/GTS / forward

Method forward

code/NNModel/attention_module.py:38–57  ·  view source on GitHub ↗

Implements Figure 2

(self, query, key, value, mask=None)

Source from the content-addressed store, hash-verified

36 self.dropout = torch.nn.Dropout(p=dropout)
37
38 def forward(self, query, key, value, mask=None):
39 "Implements Figure 2"
40 if mask is not None:
41 # Same mask applied to all h heads.
42 mask = mask.unsqueeze(1)
43 nbatches = query.size(0)
44
45 # 1) Do all the linear projections in batch from d_model => h x d_k
46 query, key, value = \
47 [l(x).view(nbatches, -1, self.h, self.d_k).transpose(1, 2)
48 for l, x in zip(self.linears, (query, key, value))]
49
50 # 2) Apply attention on all the projected vectors in batch.
51 x, self.attn = attention(query, key, value, mask=mask,
52 dropout=self.dropout)
53
54 # 3) "Concat" using a view and apply a final linear.
55 x = x.transpose(1, 2).contiguous() \
56 .view(nbatches, -1, self.h * self.d_k)
57 return self.linears[-1](x)
58
59
60class SelfAttention(torch.nn.Module):

Callers

nothing calls this directly

Calls 1

attentionFunction · 0.85

Tested by

no test coverage detected