| 28 | |
| 29 | |
| 30 | class LearnedPositionalEncoding(nn.Module): |
| 31 | |
| 32 | def __init__(self, d_model, dropout=0.1, max_len=5000): |
| 33 | super(LearnedPositionalEncoding, self).__init__() |
| 34 | self.dropout = nn.Dropout(p=dropout) |
| 35 | self.pe = nn.Parameter(torch.randn(max_len, 1, d_model)) |
| 36 | |
| 37 | def forward(self, x): |
| 38 | x = x + self.pe[:x.shape[0]] |
| 39 | return self.dropout(x) |
| 40 | |
| 41 | |
| 42 | def timestep_embedding(timesteps, dim, max_period=10000): |