https://github.com/evelinehong/Transformer_Relative_Position_PyTorch/blob/master/relative_position.py
| 18 | |
| 19 | |
| 20 | class RelativePosition(nn.Module): |
| 21 | """ https://github.com/evelinehong/Transformer_Relative_Position_PyTorch/blob/master/relative_position.py """ |
| 22 | |
| 23 | def __init__(self, num_units, max_relative_position): |
| 24 | super().__init__() |
| 25 | self.num_units = num_units |
| 26 | self.max_relative_position = max_relative_position |
| 27 | self.embeddings_table = nn.Parameter(torch.Tensor(max_relative_position * 2 + 1, num_units)) |
| 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): |