| 6 | |
| 7 | |
| 8 | class TransformerEmbedding(nn.Module): |
| 9 | def __init__(self, vocab_size, embed_size, max_len): |
| 10 | super(TransformerEmbedding, self).__init__() |
| 11 | self.embed_size = embed_size |
| 12 | |
| 13 | # 词嵌入层 |
| 14 | self.word_embedding = nn.Embedding(vocab_size, embed_size) |
| 15 | |
| 16 | # 位置编码 |
| 17 | self.position_encoding = self.create_position_encoding(max_len, embed_size) |
| 18 | |
| 19 | def create_position_encoding(self, max_len, embed_size): |
| 20 | position_encoding = torch.zeros(max_len, embed_size) |
| 21 | for pos in range(max_len): |
| 22 | for i in range(0, embed_size, 2): |
| 23 | position_encoding[pos, i] = math.sin(pos / (10000 ** ((2 * i)/embed_size))) |
| 24 | if i+1 < embed_size: |
| 25 | position_encoding[pos, i+1] = math.cos(pos / (10000 ** ((2 * i)/embed_size))) |
| 26 | return position_encoding.unsqueeze(0) |
| 27 | |
| 28 | def forward(self, x): |
| 29 | # x是输入的词索引序列,形状为(batch_size, seq_len) |
| 30 | |
| 31 | seq_len = x.size(1) |
| 32 | |
| 33 | # 获取词嵌入 |
| 34 | word_embeddings = self.word_embedding(x) |
| 35 | |
| 36 | # 添加位置编码 |
| 37 | position_embeddings = self.position_encoding[:, :seq_len, :].to(x.device) |
| 38 | |
| 39 | embeddings = word_embeddings + position_embeddings |
| 40 | return embeddings |
| 41 | |
| 42 | |
| 43 | if __name__ == "__main__": |