MCPcopy Create free account
hub / github.com/pytorch/examples / Attention

Class Attention

distributed/FSDP2/model.py:18–57  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

16
17
18class Attention(nn.Module):
19 def __init__(self, args: ModelArgs):
20 super().__init__()
21 assert args.dim % args.n_heads == 0
22 self.head_dim = args.dim // args.n_heads
23 self.n_heads = args.n_heads
24 self.dropout_p = args.dropout_p
25 self.resid_dropout = nn.Dropout(args.dropout_p)
26
27 self.wq = nn.Linear(args.dim, args.dim, bias=False)
28 self.wk = nn.Linear(args.dim, args.dim, bias=False)
29 self.wv = nn.Linear(args.dim, args.dim, bias=False)
30 self.wo = nn.Linear(args.dim, args.dim, bias=False)
31
32 def forward(self, x):
33 bsz, seq_len, _ = x.size()
34 queries, keys, values = self.wq(x), self.wk(x), self.wv(x)
35 queries = queries.view(bsz, seq_len, self.n_heads, self.head_dim)
36 keys = keys.view(bsz, seq_len, self.n_heads, self.head_dim)
37 values = values.view(bsz, seq_len, self.n_heads, self.head_dim)
38
39 queries = queries.transpose(1, 2) # (bsz, n_heads, seq_len, head_dim)
40 keys = keys.transpose(1, 2) # (bsz, n_heads, seq_len, head_dim)
41 values = values.transpose(1, 2) # (bsz, n_heads, seq_len, head_dim)
42
43 output = F.scaled_dot_product_attention(
44 queries,
45 keys,
46 values,
47 None,
48 self.dropout_p if self.training else 0,
49 )
50 output = output.transpose(1, 2).contiguous().view(bsz, seq_len, -1)
51 return self.resid_dropout(self.wo(output))
52
53 def reset_parameters(self):
54 self.wq.reset_parameters()
55 self.wk.reset_parameters()
56 self.wv.reset_parameters()
57 self.wo.reset_parameters()
58
59
60class FeedForward(nn.Module):

Callers 1

__init__Method · 0.70

Calls

no outgoing calls

Tested by

no test coverage detected