| 3 | |
| 4 | |
| 5 | class PositionalEncoder(nn.Module): |
| 6 | def __init__(self, d, T=730, repeat=None, offset=0): |
| 7 | super(PositionalEncoder, self).__init__() |
| 8 | self.d = d |
| 9 | self.T = T |
| 10 | self.repeat = repeat |
| 11 | self.denom = torch.pow( |
| 12 | T, 2 * torch.div(torch.arange(offset, offset + d).float(), 2, rounding_mode='floor') / (d+offset) |
| 13 | ) |
| 14 | self.updated_location = False |
| 15 | |
| 16 | def forward(self, batch_positions): |
| 17 | if not self.updated_location: |
| 18 | self.denom = self.denom.to(batch_positions.device) |
| 19 | self.updated_location = True |
| 20 | sinusoid_table = ( |
| 21 | batch_positions[:, :, None] / self.denom[None, None, :] |
| 22 | ) # B x T x C |
| 23 | sinusoid_table[:, :, 0::2] = torch.sin(sinusoid_table[:, :, 0::2]) # dim 2i |
| 24 | sinusoid_table[:, :, 1::2] = torch.cos(sinusoid_table[:, :, 1::2]) # dim 2i+1 |
| 25 | |
| 26 | if self.repeat is not None: |
| 27 | sinusoid_table = torch.cat( |
| 28 | [sinusoid_table for _ in range(self.repeat)], dim=-1 |
| 29 | ) |
| 30 | |
| 31 | return sinusoid_table |