MCPcopy Create free account
hub / github.com/Sirwenhao/Deep-Learning-Notes / TransformerEmbedding

Class TransformerEmbedding

NLP/Transformer/Embedding_layer.py:8–40  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

6
7
8class 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
43if __name__ == "__main__":

Callers 1

Embedding_layer.pyFile · 0.85

Calls

no outgoing calls

Tested by

no test coverage detected