| 5 | from torch import nn |
| 6 | |
| 7 | class PositionalEncoding(nn.Module): |
| 8 | def __init__( |
| 9 | self, |
| 10 | emb_size, |
| 11 | dropout, |
| 12 | maxlen=5000 |
| 13 | ): |
| 14 | super(PositionalEncoding, self).__init__() |
| 15 | den = torch.exp(- torch.arange(0, emb_size, 2)* math.log(10000) / emb_size) |
| 16 | pos = torch.arange(0, maxlen).reshape(maxlen, 1) |
| 17 | pos_embedding = torch.zeros((maxlen, emb_size)) |
| 18 | pos_embedding[:, 0::2] = torch.sin(pos * den) |
| 19 | pos_embedding[:, 1::2] = torch.cos(pos * den) |
| 20 | pos_embedding = pos_embedding.unsqueeze(-2) |
| 21 | |
| 22 | self.dropout = nn.Dropout(dropout) |
| 23 | self.register_buffer('pos_embedding', pos_embedding) |
| 24 | |
| 25 | def forward(self, token_embedding): |
| 26 | return self.dropout(token_embedding + self.pos_embedding[:token_embedding.size(0), :]) |
| 27 | |
| 28 | class Translator(nn.Module): |
| 29 | def __init__( |