(self, length_q, length_k)
| 28 | nn.init.xavier_uniform_(self.embeddings_table) |
| 29 | |
| 30 | def forward(self, length_q, length_k): |
| 31 | device = self.embeddings_table.device |
| 32 | range_vec_q = torch.arange(length_q, device=device) |
| 33 | range_vec_k = torch.arange(length_k, device=device) |
| 34 | distance_mat = range_vec_k[None, :] - range_vec_q[:, None] |
| 35 | distance_mat_clipped = torch.clamp(distance_mat, -self.max_relative_position, self.max_relative_position) |
| 36 | final_mat = distance_mat_clipped + self.max_relative_position |
| 37 | # final_mat = th.LongTensor(final_mat).to(self.embeddings_table.device) |
| 38 | # final_mat = th.tensor(final_mat, device=self.embeddings_table.device, dtype=torch.long) |
| 39 | final_mat = final_mat.long() |
| 40 | embeddings = self.embeddings_table[final_mat] |
| 41 | return embeddings |
| 42 | |
| 43 | |
| 44 | class CrossAttention(nn.Module): |
nothing calls this directly
no outgoing calls
no test coverage detected