(self, x)
| 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() |
nothing calls this directly
no outgoing calls
no test coverage detected