CBoW-TF

训练 CBoW 模型#

本笔记本是 AI for Beginners Curriculum 的一部分

在这个例子中,我们将学习如何训练 CBoW 语言模型,以获得我们自己的 Word2Vec 嵌入空间。我们将使用 AG News 数据集作为文本来源。

In [30]:
from tensorflow import keras
import tensorflow as tf
import tensorflow_datasets as tfds
import numpy as np

我们将从加载数据集开始:

In [1]:
ds_train, ds_test = tfds.load('ag_news_subset').values()

CBoW 模型#

CBoW 学习通过 $2N$ 个相邻的单词来预测一个单词。例如,当 $N=1$ 时,我们可以从句子 I like to train networks 中得到以下配对:(like, I)、(I, like)、(to, like)、(like, to)、(train, to)、(to, train)、(networks, train)、(train, networks)。这里,第一个单词是作为输入的相邻单词,第二个单词是我们要预测的单词。

为了构建一个预测下一个单词的网络,我们需要将相邻单词作为输入,并输出单词编号。CBoW 网络的架构如下:

  • 输入单词会通过嵌入层(embedding layer)。这个嵌入层将成为我们的 Word2Vec 嵌入,因此我们会将其单独定义为 embedder 变量。在这个例子中,我们将使用嵌入维度为 30,尽管你可能想尝试更高的维度(真实的 Word2Vec 通常有 300 维)。
  • 嵌入向量随后会传递到一个全连接层(dense layer),用于预测输出单词。因此,这一层会有 vocab_size 个神经元。

Keras 中的嵌入层会自动将数值输入转换为独热编码(one-hot encoding),因此我们不需要单独对输入单词进行独热编码。我们通过指定 input_length=1 来表明我们只需要输入序列中的一个单词——通常嵌入层是为处理更长的序列设计的。

对于输出,如果我们使用 sparse_categorical_crossentropy 作为损失函数,我们只需要提供单词编号作为期望结果,而不需要进行独热编码。

我们将 vocab_size 设置为 5000,以减少计算量。同时,我们还会定义一个稍后会用到的向量化工具(vectorizer)。

In [68]:
vocab_size = 5000

vectorizer = keras.layers.experimental.preprocessing.TextVectorization(max_tokens=vocab_size,input_shape=(1,))
embedder = keras.layers.Embedding(vocab_size,30,input_length=1)

model = keras.Sequential([
    embedder,
    keras.layers.Dense(vocab_size,activation='softmax')
])

model.summary()
Model: "sequential_1"
_________________________________________________________________
 Layer (type)                Output Shape              Param #   
=================================================================
 embedding_1 (Embedding)     (None, 1, 30)             150000    
                                                                 
 dense_1 (Dense)             (None, 1, 5000)           155000    
                                                                 
=================================================================
Total params: 305,000
Trainable params: 305,000
Non-trainable params: 0
_________________________________________________________________

让我们初始化向量化器并获取词汇表:

In [69]:
def extract_text(x):
    return x['title']+' '+x['description']

vectorizer.adapt(ds_train.take(500).map(extract_text))
vocab = vectorizer.get_vocabulary()

准备训练数据#

现在让我们编写一个主函数,用于从文本中计算CBoW词对。这个函数将允许我们指定窗口大小,并返回一组词对——输入词和输出词。请注意,这个函数既可以用于单词,也可以用于向量/张量——这将使我们能够在传递给to_cbow函数之前对文本进行编码。

In [70]:
def to_cbow(sent,window_size=2):
    res = []
    for i,x in enumerate(sent):
        for j in range(max(0,i-window_size),min(i+window_size+1,len(sent))):
            if i!=j:
                res.append([sent[j],x])
    return res

print(to_cbow(['I','like','to','train','networks']))
print(to_cbow(vectorizer('I like to train networks')))
[['like', 'I'], ['to', 'I'], ['I', 'like'], ['to', 'like'], ['train', 'like'], ['I', 'to'], ['like', 'to'], ['train', 'to'], ['networks', 'to'], ['like', 'train'], ['to', 'train'], ['networks', 'train'], ['to', 'networks'], ['train', 'networks']]
[[<tf.Tensor: shape=(), dtype=int64, numpy=376>, <tf.Tensor: shape=(), dtype=int64, numpy=771>], [<tf.Tensor: shape=(), dtype=int64, numpy=3>, <tf.Tensor: shape=(), dtype=int64, numpy=771>], [<tf.Tensor: shape=(), dtype=int64, numpy=771>, <tf.Tensor: shape=(), dtype=int64, numpy=376>], [<tf.Tensor: shape=(), dtype=int64, numpy=3>, <tf.Tensor: shape=(), dtype=int64, numpy=376>], [<tf.Tensor: shape=(), dtype=int64, numpy=1>, <tf.Tensor: shape=(), dtype=int64, numpy=376>], [<tf.Tensor: shape=(), dtype=int64, numpy=771>, <tf.Tensor: shape=(), dtype=int64, numpy=3>], [<tf.Tensor: shape=(), dtype=int64, numpy=376>, <tf.Tensor: shape=(), dtype=int64, numpy=3>], [<tf.Tensor: shape=(), dtype=int64, numpy=1>, <tf.Tensor: shape=(), dtype=int64, numpy=3>], [<tf.Tensor: shape=(), dtype=int64, numpy=1045>, <tf.Tensor: shape=(), dtype=int64, numpy=3>], [<tf.Tensor: shape=(), dtype=int64, numpy=376>, <tf.Tensor: shape=(), dtype=int64, numpy=1>], [<tf.Tensor: shape=(), dtype=int64, numpy=3>, <tf.Tensor: shape=(), dtype=int64, numpy=1>], [<tf.Tensor: shape=(), dtype=int64, numpy=1045>, <tf.Tensor: shape=(), dtype=int64, numpy=1>], [<tf.Tensor: shape=(), dtype=int64, numpy=3>, <tf.Tensor: shape=(), dtype=int64, numpy=1045>], [<tf.Tensor: shape=(), dtype=int64, numpy=1>, <tf.Tensor: shape=(), dtype=int64, numpy=1045>]]

让我们准备训练数据集。我们将遍历所有新闻,调用 to_cbow 获取单词对列表,并将这些对添加到 XY 中。为了节省时间,我们只考虑前 10k 条新闻——如果你有更多时间等待,并希望获得更好的嵌入,可以轻松去掉这个限制 :)

In [100]:
X = []
Y = []
for i,x in zip(range(10000),ds_train.map(extract_text).as_numpy_iterator()):
    for w1, w2 in to_cbow(vectorizer(x),window_size=1):
        X.append(tf.expand_dims(w1,0))
        Y.append(tf.expand_dims(w2,0))

我们还将把这些数据转换为一个数据集,并将其分批用于训练:

In [101]:
ds = tf.data.Dataset.from_tensor_slices((X,Y)).batch(256)

现在让我们进行实际训练。我们将使用SGD优化器,并设置较高的学习率。你也可以尝试使用其他优化器,比如Adam。我们将从训练200个周期开始——如果你希望获得更低的损失,可以重新运行此单元格。

In [102]:
model.compile(optimizer=keras.optimizers.SGD(lr=0.1),loss='sparse_categorical_crossentropy')
model.fit(ds,epochs=200)
Epoch 1/200
/usr/local/lib/python3.7/dist-packages/keras/optimizer_v2/gradient_descent.py:102: UserWarning: The `lr` argument is deprecated, use `learning_rate` instead.
  super(SGD, self).__init__(name, **kwargs)
2156/2156 [==============================] - 7s 3ms/step - loss: 5.6134
Epoch 2/200
2156/2156 [==============================] - 7s 3ms/step - loss: 5.5431
Epoch 3/200
2156/2156 [==============================] - 7s 3ms/step - loss: 5.5029
Epoch 4/200
2156/2156 [==============================] - 7s 3ms/step - loss: 5.4754
Epoch 5/200
2156/2156 [==============================] - 7s 3ms/step - loss: 5.4548
Epoch 6/200
2156/2156 [==============================] - 7s 3ms/step - loss: 5.4382
Epoch 7/200
2156/2156 [==============================] - 7s 3ms/step - loss: 5.4243
Epoch 8/200
2156/2156 [==============================] - 7s 3ms/step - loss: 5.4123
Epoch 9/200
2156/2156 [==============================] - 7s 3ms/step - loss: 5.4019
Epoch 10/200
2156/2156 [==============================] - 7s 3ms/step - loss: 5.3926
Epoch 11/200
2156/2156 [==============================] - 7s 3ms/step - loss: 5.3843
Epoch 12/200
2156/2156 [==============================] - 7s 3ms/step - loss: 5.3767
Epoch 13/200
2156/2156 [==============================] - 7s 3ms/step - loss: 5.3697
Epoch 14/200
2156/2156 [==============================] - 7s 3ms/step - loss: 5.3632
Epoch 15/200
2156/2156 [==============================] - 7s 3ms/step - loss: 5.3571
Epoch 16/200
2156/2156 [==============================] - 7s 3ms/step - loss: 5.3513
Epoch 17/200
2156/2156 [==============================] - 7s 3ms/step - loss: 5.3459
Epoch 18/200
2156/2156 [==============================] - 7s 3ms/step - loss: 5.3408
Epoch 19/200
2156/2156 [==============================] - 7s 3ms/step - loss: 5.3359
Epoch 20/200
2156/2156 [==============================] - 7s 3ms/step - loss: 5.3312
Epoch 21/200
2156/2156 [==============================] - 7s 3ms/step - loss: 5.3266
Epoch 22/200
2156/2156 [==============================] - 7s 3ms/step - loss: 5.3223
Epoch 23/200
2156/2156 [==============================] - 7s 3ms/step - loss: 5.3181
Epoch 24/200
2156/2156 [==============================] - 7s 3ms/step - loss: 5.3140
Epoch 25/200
2156/2156 [==============================] - 7s 3ms/step - loss: 5.3101
Epoch 26/200
2156/2156 [==============================] - 7s 3ms/step - loss: 5.3062
Epoch 27/200
2156/2156 [==============================] - 7s 3ms/step - loss: 5.3025
Epoch 28/200
2156/2156 [==============================] - 7s 3ms/step - loss: 5.2989
Epoch 29/200
2156/2156 [==============================] - 7s 3ms/step - loss: 5.2953
Epoch 30/200
2156/2156 [==============================] - 7s 3ms/step - loss: 5.2919
Epoch 31/200
2156/2156 [==============================] - 7s 3ms/step - loss: 5.2885
Epoch 32/200
2156/2156 [==============================] - 7s 3ms/step - loss: 5.2852
Epoch 33/200
2156/2156 [==============================] - 7s 3ms/step - loss: 5.2819
Epoch 34/200
2156/2156 [==============================] - 7s 3ms/step - loss: 5.2787
Epoch 35/200
2156/2156 [==============================] - 7s 3ms/step - loss: 5.2756
Epoch 36/200
2156/2156 [==============================] - 7s 3ms/step - loss: 5.2725
Epoch 37/200
2156/2156 [==============================] - 7s 3ms/step - loss: 5.2695
Epoch 38/200
2156/2156 [==============================] - 7s 3ms/step - loss: 5.2665
Epoch 39/200
2156/2156 [==============================] - 7s 3ms/step - loss: 5.2636
Epoch 40/200
2156/2156 [==============================] - 7s 3ms/step - loss: 5.2607
Epoch 41/200
2156/2156 [==============================] - 7s 3ms/step - loss: 5.2578
Epoch 42/200
2156/2156 [==============================] - 7s 3ms/step - loss: 5.2550
Epoch 43/200
2156/2156 [==============================] - 7s 3ms/step - loss: 5.2523
Epoch 44/200
2156/2156 [==============================] - 7s 3ms/step - loss: 5.2495
Epoch 45/200
2156/2156 [==============================] - 7s 3ms/step - loss: 5.2468
Epoch 46/200
2156/2156 [==============================] - 7s 3ms/step - loss: 5.2442
Epoch 47/200
2156/2156 [==============================] - 7s 3ms/step - loss: 5.2416
Epoch 48/200
2156/2156 [==============================] - 7s 3ms/step - loss: 5.2390
Epoch 49/200
2156/2156 [==============================] - 7s 3ms/step - loss: 5.2364
Epoch 50/200
2156/2156 [==============================] - 7s 3ms/step - loss: 5.2339
Epoch 51/200
2156/2156 [==============================] - 7s 3ms/step - loss: 5.2314
Epoch 52/200
2156/2156 [==============================] - 7s 3ms/step - loss: 5.2290
Epoch 53/200
2156/2156 [==============================] - 7s 3ms/step - loss: 5.2266
Epoch 54/200
2156/2156 [==============================] - 7s 3ms/step - loss: 5.2242
Epoch 55/200
2156/2156 [==============================] - 7s 3ms/step - loss: 5.2218
Epoch 56/200
2156/2156 [==============================] - 7s 3ms/step - loss: 5.2195
Epoch 57/200
2156/2156 [==============================] - 7s 3ms/step - loss: 5.2172
Epoch 58/200
2156/2156 [==============================] - 7s 3ms/step - loss: 5.2149
Epoch 59/200
2156/2156 [==============================] - 7s 3ms/step - loss: 5.2126
Epoch 60/200
2156/2156 [==============================] - 7s 3ms/step - loss: 5.2104
Epoch 61/200
2156/2156 [==============================] - 7s 3ms/step - loss: 5.2082
Epoch 62/200
2156/2156 [==============================] - 7s 3ms/step - loss: 5.2060
Epoch 63/200
2156/2156 [==============================] - 7s 3ms/step - loss: 5.2038
Epoch 64/200
2156/2156 [==============================] - 7s 3ms/step - loss: 5.2017
Epoch 65/200
2156/2156 [==============================] - 7s 3ms/step - loss: 5.1996
Epoch 66/200
2156/2156 [==============================] - 7s 3ms/step - loss: 5.1975
Epoch 67/200
2156/2156 [==============================] - 7s 3ms/step - loss: 5.1954
Epoch 68/200
2156/2156 [==============================] - 7s 3ms/step - loss: 5.1933
Epoch 69/200
2156/2156 [==============================] - 7s 3ms/step - loss: 5.1913
Epoch 70/200
2156/2156 [==============================] - 7s 3ms/step - loss: 5.1893
Epoch 71/200
2156/2156 [==============================] - 7s 3ms/step - loss: 5.1873
Epoch 72/200
2156/2156 [==============================] - 7s 3ms/step - loss: 5.1853
Epoch 73/200
2156/2156 [==============================] - 7s 3ms/step - loss: 5.1833
Epoch 74/200
2156/2156 [==============================] - 7s 3ms/step - loss: 5.1814
Epoch 75/200
2156/2156 [==============================] - 7s 3ms/step - loss: 5.1795
Epoch 76/200
2156/2156 [==============================] - 7s 3ms/step - loss: 5.1775
Epoch 77/200
2156/2156 [==============================] - 7s 3ms/step - loss: 5.1756
Epoch 78/200
2156/2156 [==============================] - 7s 3ms/step - loss: 5.1737
Epoch 79/200
2156/2156 [==============================] - 7s 3ms/step - loss: 5.1719
Epoch 80/200
2156/2156 [==============================] - 7s 3ms/step - loss: 5.1700
Epoch 81/200
2156/2156 [==============================] - 7s 3ms/step - loss: 5.1682
Epoch 82/200
2156/2156 [==============================] - 7s 3ms/step - loss: 5.1663
Epoch 83/200
2156/2156 [==============================] - 7s 3ms/step - loss: 5.1645
Epoch 84/200
2156/2156 [==============================] - 7s 3ms/step - loss: 5.1627
Epoch 85/200
2156/2156 [==============================] - 7s 3ms/step - loss: 5.1609
Epoch 86/200
2156/2156 [==============================] - 7s 3ms/step - loss: 5.1592
Epoch 87/200
2156/2156 [==============================] - 7s 3ms/step - loss: 5.1574
Epoch 88/200
2156/2156 [==============================] - 7s 3ms/step - loss: 5.1557
Epoch 89/200
2156/2156 [==============================] - 7s 3ms/step - loss: 5.1539
Epoch 90/200
2156/2156 [==============================] - 7s 3ms/step - loss: 5.1522
Epoch 91/200
2156/2156 [==============================] - 7s 3ms/step - loss: 5.1505
Epoch 92/200
2156/2156 [==============================] - 7s 3ms/step - loss: 5.1488
Epoch 93/200
2156/2156 [==============================] - 7s 3ms/step - loss: 5.1471
Epoch 94/200
2156/2156 [==============================] - 7s 3ms/step - loss: 5.1454
Epoch 95/200
2156/2156 [==============================] - 7s 3ms/step - loss: 5.1438
Epoch 96/200
2156/2156 [==============================] - 7s 3ms/step - loss: 5.1421
Epoch 97/200
2156/2156 [==============================] - 7s 3ms/step - loss: 5.1405
Epoch 98/200
2156/2156 [==============================] - 7s 3ms/step - loss: 5.1389
Epoch 99/200
2156/2156 [==============================] - 7s 3ms/step - loss: 5.1372
Epoch 100/200
2156/2156 [==============================] - 7s 3ms/step - loss: 5.1356
Epoch 101/200
2156/2156 [==============================] - 7s 3ms/step - loss: 5.1341
Epoch 102/200
2156/2156 [==============================] - 7s 3ms/step - loss: 5.1325
Epoch 103/200
2156/2156 [==============================] - 7s 3ms/step - loss: 5.1309
Epoch 104/200
2156/2156 [==============================] - 7s 3ms/step - loss: 5.1293
Epoch 105/200
2156/2156 [==============================] - 7s 3ms/step - loss: 5.1278
Epoch 106/200
2156/2156 [==============================] - 7s 3ms/step - loss: 5.1263
Epoch 107/200
2156/2156 [==============================] - 7s 3ms/step - loss: 5.1247
Epoch 108/200
2156/2156 [==============================] - 7s 3ms/step - loss: 5.1232
Epoch 109/200
2156/2156 [==============================] - 7s 3ms/step - loss: 5.1217
Epoch 110/200
2156/2156 [==============================] - 7s 3ms/step - loss: 5.1202
Epoch 111/200
2156/2156 [==============================] - 7s 3ms/step - loss: 5.1187
Epoch 112/200
2156/2156 [==============================] - 7s 3ms/step - loss: 5.1173
Epoch 113/200
2156/2156 [==============================] - 7s 3ms/step - loss: 5.1158
Epoch 114/200
2156/2156 [==============================] - 7s 3ms/step - loss: 5.1144
Epoch 115/200
2156/2156 [==============================] - 7s 3ms/step - loss: 5.1129
Epoch 116/200
2156/2156 [==============================] - 7s 3ms/step - loss: 5.1115
Epoch 117/200
2156/2156 [==============================] - 7s 3ms/step - loss: 5.1101
Epoch 118/200
2156/2156 [==============================] - 7s 3ms/step - loss: 5.1086
Epoch 119/200
2156/2156 [==============================] - 7s 3ms/step - loss: 5.1072
Epoch 120/200
2156/2156 [==============================] - 7s 3ms/step - loss: 5.1058
Epoch 121/200
2156/2156 [==============================] - 7s 3ms/step - loss: 5.1045
Epoch 122/200
2156/2156 [==============================] - 7s 3ms/step - loss: 5.1031
Epoch 123/200
2156/2156 [==============================] - 7s 3ms/step - loss: 5.1017
Epoch 124/200
2156/2156 [==============================] - 7s 3ms/step - loss: 5.1004
Epoch 125/200
2156/2156 [==============================] - 7s 3ms/step - loss: 5.0990
Epoch 126/200
2156/2156 [==============================] - 7s 3ms/step - loss: 5.0977
Epoch 127/200
2156/2156 [==============================] - 7s 3ms/step - loss: 5.0963
Epoch 128/200
2156/2156 [==============================] - 7s 3ms/step - loss: 5.0950
Epoch 129/200
2156/2156 [==============================] - 7s 3ms/step - loss: 5.0937
Epoch 130/200
2156/2156 [==============================] - 7s 3ms/step - loss: 5.0924
Epoch 131/200
2156/2156 [==============================] - 7s 3ms/step - loss: 5.0911
Epoch 132/200
2156/2156 [==============================] - 7s 3ms/step - loss: 5.0898
Epoch 133/200
2156/2156 [==============================] - 7s 3ms/step - loss: 5.0885
Epoch 134/200
2156/2156 [==============================] - 7s 3ms/step - loss: 5.0873
Epoch 135/200
2156/2156 [==============================] - 7s 3ms/step - loss: 5.0860
Epoch 136/200
2156/2156 [==============================] - 7s 3ms/step - loss: 5.0848
Epoch 137/200
2156/2156 [==============================] - 7s 3ms/step - loss: 5.0835
Epoch 138/200
2156/2156 [==============================] - 7s 3ms/step - loss: 5.0823
Epoch 139/200
2156/2156 [==============================] - 7s 3ms/step - loss: 5.0810
Epoch 140/200
2156/2156 [==============================] - 7s 3ms/step - loss: 5.0798
Epoch 141/200
2156/2156 [==============================] - 7s 3ms/step - loss: 5.0786
Epoch 142/200
2156/2156 [==============================] - 7s 3ms/step - loss: 5.0774
Epoch 143/200
2156/2156 [==============================] - 7s 3ms/step - loss: 5.0762
Epoch 144/200
2156/2156 [==============================] - 7s 3ms/step - loss: 5.0750
Epoch 145/200
2156/2156 [==============================] - 8s 4ms/step - loss: 5.0739
Epoch 146/200
2156/2156 [==============================] - 7s 3ms/step - loss: 5.0727
Epoch 147/200
2156/2156 [==============================] - 7s 3ms/step - loss: 5.0715
Epoch 148/200
2156/2156 [==============================] - 7s 3ms/step - loss: 5.0704
Epoch 149/200
2156/2156 [==============================] - 7s 3ms/step - loss: 5.0692
Epoch 150/200
2156/2156 [==============================] - 7s 3ms/step - loss: 5.0681
Epoch 151/200
2156/2156 [==============================] - 7s 3ms/step - loss: 5.0670
Epoch 152/200
2156/2156 [==============================] - 7s 3ms/step - loss: 5.0658
Epoch 153/200
2156/2156 [==============================] - 7s 3ms/step - loss: 5.0647
Epoch 154/200
2156/2156 [==============================] - 7s 3ms/step - loss: 5.0636
Epoch 155/200
2156/2156 [==============================] - 7s 3ms/step - loss: 5.0625
Epoch 156/200
2156/2156 [==============================] - 7s 3ms/step - loss: 5.0614
Epoch 157/200
2156/2156 [==============================] - 7s 3ms/step - loss: 5.0603
Epoch 158/200
2156/2156 [==============================] - 7s 3ms/step - loss: 5.0593
Epoch 159/200
2156/2156 [==============================] - 7s 3ms/step - loss: 5.0582
Epoch 160/200
2156/2156 [==============================] - 7s 3ms/step - loss: 5.0571
Epoch 161/200
2156/2156 [==============================] - 7s 3ms/step - loss: 5.0561
Epoch 162/200
2156/2156 [==============================] - 7s 3ms/step - loss: 5.0550
Epoch 163/200
2156/2156 [==============================] - 7s 3ms/step - loss: 5.0539
Epoch 164/200
2156/2156 [==============================] - 7s 3ms/step - loss: 5.0529
Epoch 165/200
2156/2156 [==============================] - 7s 3ms/step - loss: 5.0519
Epoch 166/200
2156/2156 [==============================] - 7s 3ms/step - loss: 5.0508
Epoch 167/200
2156/2156 [==============================] - 7s 3ms/step - loss: 5.0498
Epoch 168/200
2156/2156 [==============================] - 7s 3ms/step - loss: 5.0488
Epoch 169/200
2156/2156 [==============================] - 7s 3ms/step - loss: 5.0478
Epoch 170/200
2156/2156 [==============================] - 7s 3ms/step - loss: 5.0468
Epoch 171/200
2156/2156 [==============================] - 7s 3ms/step - loss: 5.0458
Epoch 172/200
2156/2156 [==============================] - 7s 3ms/step - loss: 5.0448
Epoch 173/200
2156/2156 [==============================] - 7s 3ms/step - loss: 5.0438
Epoch 174/200
2156/2156 [==============================] - 7s 3ms/step - loss: 5.0428
Epoch 175/200
2156/2156 [==============================] - 7s 3ms/step - loss: 5.0418
Epoch 176/200
2156/2156 [==============================] - 7s 3ms/step - loss: 5.0409
Epoch 177/200
2156/2156 [==============================] - 7s 3ms/step - loss: 5.0399
Epoch 178/200
2156/2156 [==============================] - 7s 3ms/step - loss: 5.0389
Epoch 179/200
2156/2156 [==============================] - 7s 3ms/step - loss: 5.0380
Epoch 180/200
2156/2156 [==============================] - 7s 3ms/step - loss: 5.0370
Epoch 181/200
2156/2156 [==============================] - 7s 3ms/step - loss: 5.0361
Epoch 182/200
2156/2156 [==============================] - 7s 3ms/step - loss: 5.0351
Epoch 183/200
2156/2156 [==============================] - 7s 3ms/step - loss: 5.0342
Epoch 184/200
2156/2156 [==============================] - 7s 3ms/step - loss: 5.0333
Epoch 185/200
2156/2156 [==============================] - 7s 3ms/step - loss: 5.0323
Epoch 186/200
2156/2156 [==============================] - 7s 3ms/step - loss: 5.0314
Epoch 187/200
2156/2156 [==============================] - 7s 3ms/step - loss: 5.0305
Epoch 188/200
2156/2156 [==============================] - 7s 3ms/step - loss: 5.0296
Epoch 189/200
2156/2156 [==============================] - 7s 3ms/step - loss: 5.0287
Epoch 190/200
2156/2156 [==============================] - 7s 3ms/step - loss: 5.0278
Epoch 191/200
2156/2156 [==============================] - 7s 3ms/step - loss: 5.0269
Epoch 192/200
2156/2156 [==============================] - 7s 3ms/step - loss: 5.0260
Epoch 193/200
2156/2156 [==============================] - 7s 3ms/step - loss: 5.0251
Epoch 194/200
2156/2156 [==============================] - 7s 3ms/step - loss: 5.0242
Epoch 195/200
2156/2156 [==============================] - 7s 3ms/step - loss: 5.0233
Epoch 196/200
2156/2156 [==============================] - 7s 3ms/step - loss: 5.0225
Epoch 197/200
2156/2156 [==============================] - 7s 3ms/step - loss: 5.0216
Epoch 198/200
2156/2156 [==============================] - 7s 3ms/step - loss: 5.0207
Epoch 199/200
2156/2156 [==============================] - 7s 3ms/step - loss: 5.0199
Epoch 200/200
2156/2156 [==============================] - 7s 3ms/step - loss: 5.0190
<keras.callbacks.History at 0x7ff7e52572d0>

尝试使用 Word2Vec#

为了使用 Word2Vec,让我们提取与词汇表中所有单词对应的向量:

In [103]:
vectors = embedder(vectorizer(vocab))
vectors = tf.reshape(vectors,(-1,30)) # we need reshape to get rid of extra dimension

让我们来看一个例子,单词Paris是如何被编码成一个向量的:

In [104]:
paris_vec = embedder(vectorizer('paris'))[0]
print(paris_vec)
tf.Tensor(
[-0.13308628  0.50972325  0.00344684  0.185389   -0.03176536  0.22262476
 -0.3856765  -0.6854793   0.5185803  -0.7215402  -0.16101503  0.15622072
  0.00653811 -0.14954254  0.03379822 -0.01243829  0.27907634 -0.32538188
  0.21718933  0.31112966 -0.24142407  0.15589055  0.2915561   0.19029242
  0.08425518 -0.0941902  -0.54313695 -0.24854654  0.26196313  0.18027727], shape=(30,), dtype=float32)

使用Word2Vec查找同义词是很有趣的。以下函数将返回与给定输入最接近的n个单词。为了找到它们,我们计算$|w_i - v|$的范数,其中$v$是对应于输入单词的向量,$w_i$是词汇表中第$i$个单词的编码。然后我们对数组进行排序,并使用argsort返回相应的索引,取列表的前n个元素,这些元素编码了词汇表中最接近单词的位置。

In [105]:
def close_words(x,n=5):
  vec = embedder(vectorizer(x))[0]
  top5 = np.linalg.norm(vectors-vec,axis=1).argsort()[:n]
  return [ vocab[x] for x in top5 ]

close_words('paris')
['paris', 'philippines', 'seoul', 'jakarta', 'zoo']
In [112]:
close_words('china')
['china', 'russia', 'pakistan', 'israel', 'turkey']
In [113]:
close_words('official')
['official', 'military', 'office', 'police', 'sources']

要点#

通过使用像CBoW这样的巧妙技术,我们可以训练Word2Vec模型。你也可以尝试训练skip-gram模型,该模型通过给定中心词预测邻近词,看看它的表现如何。


免责声明
本文档使用AI翻译服务 Co-op Translator 进行翻译。尽管我们努力确保翻译的准确性,但请注意,自动翻译可能包含错误或不准确之处。应以原文档的原始语言版本作为权威来源。对于关键信息,建议使用专业人工翻译。我们对因使用此翻译而引起的任何误解或误读不承担责任。