| 16 | |
| 17 | |
| 18 | class 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 | |
| 60 | class FeedForward(nn.Module): |