Pytorch的torch.nn.embedding()如何實現(xiàn)詞嵌入層
torch.nn.embedding()實現(xiàn)詞嵌入層
nn.embedding()其實是NLP中常用的詞嵌入層,在實現(xiàn)詞嵌入的過程中embedding層的權(quán)重用于隨機初始化詞的向量,該embedding層的權(quán)重參數(shù)在后續(xù)訓(xùn)練時會不斷更新調(diào)整,并被優(yōu)化。
nn.embedding:這是一個矩陣類,該開始時里面初始化了一個隨機矩陣,矩陣的長是字典的大小,寬是用來表示字典中每個元素的屬性向量,向量的維度根據(jù)你想要表示的元素的復(fù)雜度而定。
類實例化之后可以根據(jù)字典中元素的下標(biāo)來查找元素對應(yīng)的向量。
因為輸入的句子長度不一,有的長有的短。
長了截斷,不夠長補齊(我文中用’'填充,然后在nn.embedding層將其補0,也就是用它來表示無意義的詞,這樣在后面的max-pooling層也就自然而然會把其過濾掉,這樣就不用擔(dān)心他會影響識別。)
這里說一下它的用法:
nn.embedding()主要3個參數(shù)
- 第一個參數(shù)num_embeddings是指詞表大小
- 第二個參數(shù)embedding_dim是指你需要用多少維來表示一個符號
- 第三個參數(shù)pading_idx即需要用0填充的符號在詞表中的位置,如下,輸出中后面兩個’'都有被填充為了0.
import torch
import torch.nn as nn
#詞表
word_to_id = {'hello':0, '<PAD>':1,'world':2}
embeds = nn.Embedding(len(word_to_id), 4,padding_idx=word_to_id['<PAD>'])
text = 'hello world <PAD> <PAD>'
hello_idx = torch.LongTensor([word_to_id[i] for i in text.split()])
#詞嵌入得到詞向量
hello_embed = embeds(hello_idx)
print(hello_embed)
從以下輸出可以看到,每行代表句子中一個單詞的詞嵌入向量,句子中的每個單詞都有4維度,最后兩個0向量是時用來填充補齊的沒意義。
所以embedding層其實相當(dāng)于將前面用索引編碼的句子表示乘上embedding層的可訓(xùn)練權(quán)重得到的就是詞嵌入的結(jié)果
輸出:
tensor([[-1.1436, 1.4588, -1.2755, 0.0077],
[-0.9600, -1.9986, -1.1087, -0.1520],
[ 0.0000, 0.0000, 0.0000, 0.0000],
[ 0.0000, 0.0000, 0.0000, 0.0000]], grad_fn=)
你也可以使用nn.Embedding.from_pretrained()加載預(yù)訓(xùn)練好的模型,如word2vec,glove等,在訓(xùn)練的過程中也可以邊訓(xùn)練,邊更新詞向量,加快模型的收斂。
本文用的只是簡單的nn.embedding()
然后具體使用 nn.embedding() 時,寫在初始化搭建網(wǎng)絡(luò)里
如下:
class Network(nn.Module):
def __init__(self):
super(TextCNN, self).__init__(nvocab,embed)
self.filter_sizes = (2, 3, 4)
self.embed = embed
self.num_filters = 256
self.dropout = 0.5
self.num_classes = num_classes
self.n_vocab = nvocab
#通過padding_idx將<PAD>字符填充為0,因為他沒意義哦,后面max-pooling自然而然會把他過濾掉哦
self.embedding = nn.Embedding(self.n_vocab, self.embed, padding_idx=word2idx['<PAD>'])
self.convs = nn.ModuleList(
[nn.Conv2d(1, self.num_filters, (k, self.embed)) for k in self.filter_sizes])
self.dropout = nn.Dropout(self.dropout)
self.fc = nn.Linear(self.num_filters * len(self.filter_sizes), self.num_classes)
def conv_and_pool(self, x, conv):
x = F.relu(conv(x)).squeeze(3)
x = F.max_pool1d(x, x.size(2)).squeeze(2)
return x
def forward(self, x):
out = self.embedding(x)
out = out.unsqueeze(1)
out = torch.cat([self.conv_and_pool(out, conv) for conv in self.convs], 1)
out = self.dropout(out)
out = self.fc(out)
return out
總結(jié)
以上為個人經(jīng)驗,希望能給大家一個參考,也希望大家多多支持腳本之家。
相關(guān)文章
Python設(shè)計模式結(jié)構(gòu)型享元模式
這篇文章主要介紹了Python享元模式,享元模式即Flyweight Pattern,指運用共享技術(shù)有效地支持大量細粒度的對象,下面和小編一起進入文章了解更多詳細內(nèi)容吧2022-02-02
Python深度學(xué)習(xí)之實現(xiàn)卷積神經(jīng)網(wǎng)絡(luò)
今天帶大家學(xué)習(xí)如何使用Python實現(xiàn)卷積神經(jīng)網(wǎng)絡(luò),這是個很難的知識點,文中有非常詳細的介紹,對小伙伴們很有幫助,需要的朋友可以參考下2021-06-06
Pygame庫200行代碼實現(xiàn)簡易飛機大戰(zhàn)
本文主要介紹了Pygame庫200行代碼實現(xiàn)簡易飛機大戰(zhàn),文中通過示例代碼介紹的非常詳細,具有一定的參考價值,感興趣的小伙伴們可以參考一下2021-12-12

