AI-For-Beginners/translations/zh/lessons/5-NLP/15-LanguageModeling/CBoW-PyTorch.ipynb

576 lines
17 KiB
Plaintext
Raw Blame History

This file contains ambiguous Unicode characters

This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.

{
"cells": [
{
"cell_type": "markdown",
"metadata": {
"id": "NXTSugt6ieXh"
},
"source": [
"## 训练 CBoW 模型\n",
"\n",
"本笔记本是 [AI for Beginners Curriculum](http://aka.ms/ai-beginners) 的一部分\n",
"\n",
"在这个例子中,我们将学习如何训练 CBoW 语言模型,以获得我们自己的 Word2Vec 嵌入空间。我们将使用 AG News 数据集作为文本来源。\n"
]
},
{
"cell_type": "code",
"source": [
"import torch\n",
"import torchtext\n",
"import os\n",
"import collections\n",
"import builtins\n",
"import random\n",
"import numpy as np"
],
"metadata": {
"id": "q-UiiJUKaxHj"
},
"execution_count": null,
"outputs": []
},
{
"cell_type": "code",
"source": [
"device = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")"
],
"metadata": {
"id": "TFbR8CZaTZ1q"
},
"execution_count": null,
"outputs": []
},
{
"cell_type": "markdown",
"source": [
"首先,让我们加载数据集并定义分词器和词汇表。我们将 `vocab_size` 设置为 5000 以稍微限制计算量。\n"
],
"metadata": {
"id": "HIwC7lI5T-ov"
}
},
{
"cell_type": "code",
"source": [
"def load_dataset(ngrams = 1, min_freq = 1, vocab_size = 5000 , lines_cnt = 500):\n",
" tokenizer = torchtext.data.utils.get_tokenizer('basic_english')\n",
" print(\"Loading dataset...\")\n",
" test_dataset, train_dataset = torchtext.datasets.AG_NEWS(root='./data')\n",
" train_dataset = list(train_dataset)\n",
" test_dataset = list(test_dataset)\n",
" classes = ['World', 'Sports', 'Business', 'Sci/Tech']\n",
" print('Building vocab...')\n",
" counter = collections.Counter()\n",
" for i, (_, line) in enumerate(train_dataset):\n",
" counter.update(torchtext.data.utils.ngrams_iterator(tokenizer(line),ngrams=ngrams))\n",
" if i == lines_cnt:\n",
" break\n",
" vocab = torchtext.vocab.Vocab(collections.Counter(dict(counter.most_common(vocab_size))), min_freq=min_freq)\n",
" return train_dataset, test_dataset, classes, vocab, tokenizer"
],
"metadata": {
"id": "wdZuygtgiuLG"
},
"execution_count": null,
"outputs": []
},
{
"cell_type": "code",
"source": [
"train_dataset, test_dataset, _, vocab, tokenizer = load_dataset()"
],
"metadata": {
"colab": {
"base_uri": "https://localhost:8080/"
},
"id": "4d1nU1gsivGu",
"outputId": "949fe272-ae0e-49f5-c373-6703458b3a74"
},
"execution_count": null,
"outputs": [
{
"output_type": "stream",
"name": "stdout",
"text": [
"Loading dataset...\n",
"Building vocab...\n"
]
}
]
},
{
"cell_type": "code",
"source": [
"def encode(x, vocabulary, tokenizer = tokenizer):\n",
" return [vocabulary[s] for s in tokenizer(x)]"
],
"metadata": {
"id": "1XDYNhG8ToFV"
},
"execution_count": null,
"outputs": []
},
{
"cell_type": "markdown",
"metadata": {
"id": "LIlQk6_PaHVY"
},
"source": [
"## CBoW 模型\n",
"\n",
"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)。这里,第一个单词是作为输入的邻近词,第二个单词是我们要预测的目标词。\n",
"\n",
"为了构建一个预测下一个单词的网络我们需要将邻近词作为输入并输出单词编号。CBoW 网络的架构如下:\n",
"\n",
"* 输入单词会通过嵌入层处理。这个嵌入层就是我们的 Word2Vec 嵌入,因此我们会将其单独定义为 `embedder` 变量。在这个例子中,我们将使用嵌入维度为 30尽管你可能希望尝试更高的维度真实的 Word2Vec 通常使用 300。\n",
"* 嵌入向量接着会传递到一个线性层,该层用于预测输出单词。因此它包含 `vocab_size` 个神经元。\n",
"\n",
"对于输出,如果我们使用 `CrossEntropyLoss` 作为损失函数,那么我们只需要提供单词编号作为期望结果,而不需要进行独热编码。\n"
]
},
{
"cell_type": "code",
"source": [
"vocab_size = len(vocab)\n",
"\n",
"embedder = torch.nn.Embedding(num_embeddings = vocab_size, embedding_dim = 30)\n",
"model = torch.nn.Sequential(\n",
" embedder,\n",
" torch.nn.Linear(in_features = 30, out_features = vocab_size),\n",
")\n",
"\n",
"print(model)"
],
"metadata": {
"colab": {
"base_uri": "https://localhost:8080/"
},
"id": "akKTcKQKkfl2",
"outputId": "da687e3e-a8ec-4c1a-e456-ab8cd6ac7dad"
},
"execution_count": null,
"outputs": [
{
"output_type": "stream",
"name": "stdout",
"text": [
"Sequential(\n",
" (0): Embedding(5002, 30)\n",
" (1): Linear(in_features=30, out_features=5002, bias=True)\n",
")\n"
]
}
]
},
{
"cell_type": "markdown",
"metadata": {
"id": "Nud6jgGPaHVa"
},
"source": [
"## 准备训练数据\n",
"\n",
"现在让我们编写一个主函数用于从文本中计算CBoW词对。这个函数将允许我们指定窗口大小并返回一组词对——输入词和输出词。请注意这个函数既可以用于单词也可以用于向量/张量——这将使我们能够在传递给`to_cbow`函数之前对文本进行编码。\n"
]
},
{
"cell_type": "code",
"execution_count": null,
"metadata": {
"colab": {
"base_uri": "https://localhost:8080/"
},
"id": "x-dsXygOieXn",
"outputId": "c2218280-e540-40ba-9546-efe48d0d714f"
},
"outputs": [
{
"output_type": "stream",
"name": "stdout",
"text": [
"[['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']]\n",
"[[232, 172], [5, 172], [172, 232], [5, 232], [0, 232], [172, 5], [232, 5], [0, 5], [1202, 5], [232, 0], [5, 0], [1202, 0], [5, 1202], [0, 1202]]\n"
]
}
],
"source": [
"def to_cbow(sent,window_size=2):\n",
" res = []\n",
" for i,x in enumerate(sent):\n",
" for j in range(max(0,i-window_size),min(i+window_size+1,len(sent))):\n",
" if i!=j:\n",
" res.append([sent[j],x])\n",
" return res\n",
"\n",
"print(to_cbow(['I','like','to','train','networks']))\n",
"print(to_cbow(encode('I like to train networks', vocab)))"
]
},
{
"cell_type": "markdown",
"metadata": {
"id": "XVaaDLjaaHVb"
},
"source": [
"让我们准备训练数据集。我们将遍历所有新闻,调用`to_cbow`获取单词对列表,并将这些对添加到`X`和`Y`中。为了节省时间我们只考虑前1万条新闻——如果你有更多时间等待并希望获得更好的嵌入可以轻松去掉这个限制 :)\n"
]
},
{
"cell_type": "code",
"execution_count": null,
"metadata": {
"id": "54b-Gd9TieXo"
},
"outputs": [],
"source": [
"X = []\n",
"Y = []\n",
"for i, x in zip(range(10000), train_dataset):\n",
" for w1, w2 in to_cbow(encode(x[1], vocab), window_size = 5):\n",
" X.append(w1)\n",
" Y.append(w2)\n",
"\n",
"X = torch.tensor(X)\n",
"Y = torch.tensor(Y)"
]
},
{
"cell_type": "markdown",
"source": [
"我们还将把这些数据转换为一个数据集,并创建数据加载器:\n"
],
"metadata": {
"id": "cwWy0PzXWhN5"
}
},
{
"cell_type": "code",
"source": [
"class SimpleIterableDataset(torch.utils.data.IterableDataset):\n",
" def __init__(self, X, Y):\n",
" super(SimpleIterableDataset).__init__()\n",
" self.data = []\n",
" for i in range(len(X)):\n",
" self.data.append( (Y[i], X[i]) )\n",
" random.shuffle(self.data)\n",
"\n",
" def __iter__(self):\n",
" return iter(self.data)"
],
"metadata": {
"id": "mfoAcGPFZU8p"
},
"execution_count": null,
"outputs": []
},
{
"cell_type": "markdown",
"metadata": {
"id": "e4NQ_-5waHVc"
},
"source": [
"我们还将把这些数据转换为一个数据集,并创建数据加载器:\n"
]
},
{
"cell_type": "code",
"execution_count": null,
"metadata": {
"id": "AbLUcojlieXo"
},
"outputs": [],
"source": [
"ds = SimpleIterableDataset(X, Y)\n",
"dl = torch.utils.data.DataLoader(ds, batch_size = 256)"
]
},
{
"cell_type": "markdown",
"metadata": {
"id": "pKQr7sXeaHVc"
},
"source": [
"现在让我们进行实际训练。我们将使用`SGD`优化器,并设置较高的学习率。您也可以尝试使用其他优化器,例如`Adam`。我们将从训练10个周期开始——如果您希望获得更低的损失可以重新运行此单元格。\n"
]
},
{
"cell_type": "code",
"source": [
"def train_epoch(net, dataloader, lr = 0.01, optimizer = None, loss_fn = torch.nn.CrossEntropyLoss(), epochs = None, report_freq = 1):\n",
" optimizer = optimizer or torch.optim.Adam(net.parameters(), lr = lr)\n",
" loss_fn = loss_fn.to(device)\n",
" net.train()\n",
"\n",
" for i in range(epochs):\n",
" total_loss, j = 0, 0, \n",
" for labels, features in dataloader:\n",
" optimizer.zero_grad()\n",
" features, labels = features.to(device), labels.to(device)\n",
" out = net(features)\n",
" loss = loss_fn(out, labels)\n",
" loss.backward()\n",
" optimizer.step()\n",
" total_loss += loss\n",
" j += 1\n",
" if i % report_freq == 0:\n",
" print(f\"Epoch: {i+1}: loss={total_loss.item()/j}\")\n",
"\n",
" return total_loss.item()/j"
],
"metadata": {
"id": "HeeCYKr_KF1w"
},
"execution_count": null,
"outputs": []
},
{
"cell_type": "code",
"source": [
"train_epoch(net = model, dataloader = dl, optimizer = torch.optim.SGD(model.parameters(), lr = 0.1), loss_fn = torch.nn.CrossEntropyLoss(), epochs = 10)"
],
"metadata": {
"colab": {
"base_uri": "https://localhost:8080/"
},
"id": "KVgwGtDHgDlT",
"outputId": "2447833f-f0e3-4566-c33d-addbfe2f451d"
},
"execution_count": null,
"outputs": [
{
"output_type": "stream",
"name": "stdout",
"text": [
"Epoch: 1: loss=5.664632366860172\n",
"Epoch: 2: loss=5.632101973960962\n",
"Epoch: 3: loss=5.610399051405015\n",
"Epoch: 4: loss=5.594621561080262\n",
"Epoch: 5: loss=5.582538017415446\n",
"Epoch: 6: loss=5.572900234519603\n",
"Epoch: 7: loss=5.564951676341915\n",
"Epoch: 8: loss=5.558288112064614\n",
"Epoch: 9: loss=5.552576955031129\n",
"Epoch: 10: loss=5.547634165194347\n"
]
},
{
"output_type": "execute_result",
"data": {
"text/plain": [
"5.547634165194347"
]
},
"metadata": {},
"execution_count": 16
}
]
},
{
"cell_type": "markdown",
"metadata": {
"id": "W8u2qXZmaHVd"
},
"source": [
"## 尝试使用 Word2Vec\n",
"\n",
"为了使用 Word2Vec让我们提取与词汇表中所有单词对应的向量\n"
]
},
{
"cell_type": "code",
"execution_count": null,
"metadata": {
"id": "r8TatcXjkU_t"
},
"outputs": [],
"source": [
"vectors = torch.stack([embedder(torch.tensor(vocab[s])) for s in vocab.itos], 0)"
]
},
{
"cell_type": "markdown",
"metadata": {
"id": "3OcX21UOaHVd"
},
"source": [
"让我们看看,例如,单词**Paris**是如何编码成一个向量的:\n"
]
},
{
"cell_type": "code",
"execution_count": null,
"metadata": {
"colab": {
"base_uri": "https://localhost:8080/"
},
"id": "bz6tAeLzieXp",
"outputId": "5b20850e-4342-45e9-f840-cfac2b4d61d8"
},
"outputs": [
{
"output_type": "stream",
"name": "stdout",
"text": [
"tensor([-0.0915, 2.1224, -0.0281, -0.6819, 1.1219, 0.6458, -1.3704, -1.3314,\n",
" -1.1437, 0.4496, 0.2301, -0.3515, -0.8485, 1.0481, 0.4386, -0.8949,\n",
" 0.5644, 1.0939, -2.5096, 3.2949, -0.2601, -0.8640, 0.1421, -0.0804,\n",
" -0.5083, -1.0560, 0.9753, -0.5949, -1.6046, 0.5774],\n",
" grad_fn=<EmbeddingBackward>)\n"
]
}
],
"source": [
"paris_vec = embedder(torch.tensor(vocab['paris']))\n",
"print(paris_vec)"
]
},
{
"cell_type": "markdown",
"metadata": {
"id": "pHTJlaeYaHVd"
},
"source": [
"使用Word2Vec寻找同义词非常有趣。以下函数将返回与给定输入最接近的`n`个单词。为了找到它们,我们计算$|w_i - v|$的范数,其中$v$是与输入单词对应的向量,$w_i$是词汇表中第$i$个单词的编码。然后我们对数组进行排序,并使用`argsort`返回对应的索引,取列表的前`n`个元素,这些元素编码了词汇表中最接近单词的位置。\n"
]
},
{
"cell_type": "code",
"execution_count": null,
"metadata": {
"colab": {
"base_uri": "https://localhost:8080/"
},
"id": "NlZyi-_olFar",
"outputId": "b5dbb163-88c4-4d5a-eaf2-6751f700e98c"
},
"outputs": [
{
"output_type": "execute_result",
"data": {
"text/plain": [
"['microsoft', 'quoted', 'lp', 'rate', 'top']"
]
},
"metadata": {},
"execution_count": 56
}
],
"source": [
"def close_words(x, n = 5):\n",
" vec = embedder(torch.tensor(vocab[x]))\n",
" top5 = np.linalg.norm(vectors.detach().numpy() - vec.detach().numpy(), axis = 1).argsort()[:n]\n",
" return [ vocab.itos[x] for x in top5 ]\n",
"\n",
"close_words('microsoft')"
]
},
{
"cell_type": "code",
"execution_count": null,
"metadata": {
"colab": {
"base_uri": "https://localhost:8080/"
},
"id": "-dQq7xeAln0U",
"outputId": "66f768c3-c248-4bfd-ce4f-c8ffc6d0dd0d"
},
"outputs": [
{
"output_type": "execute_result",
"data": {
"text/plain": [
"['basketball', 'lot', 'sinai', 'states', 'healthdaynews']"
]
},
"metadata": {},
"execution_count": 51
}
],
"source": [
"close_words('basketball')"
]
},
{
"cell_type": "code",
"execution_count": null,
"metadata": {
"colab": {
"base_uri": "https://localhost:8080/"
},
"id": "fJXqK26b29sa",
"outputId": "78f0baba-ffd0-485a-dd87-0a12bedfd7fa"
},
"outputs": [
{
"output_type": "execute_result",
"data": {
"text/plain": [
"['funds', 'travel', 'sydney', 'japan', 'business']"
]
},
"metadata": {},
"execution_count": 77
}
],
"source": [
"close_words('funds')"
]
},
{
"cell_type": "markdown",
"metadata": {
"id": "My0VeTDd3Ji8"
},
"source": [
"## 要点\n",
"\n",
"通过使用诸如CBoW这样的巧妙技术我们可以训练Word2Vec模型。你也可以尝试训练skip-gram模型该模型通过给定中心词来预测邻近词看看它的表现如何。\n"
]
},
{
"cell_type": "markdown",
"metadata": {},
"source": [
"\n---\n\n**免责声明** \n本文档使用AI翻译服务 [Co-op Translator](https://github.com/Azure/co-op-translator) 进行翻译。尽管我们努力确保翻译的准确性,但请注意,自动翻译可能包含错误或不准确之处。应以原始语言的文档作为权威来源。对于关键信息,建议使用专业人工翻译。我们对因使用此翻译而引起的任何误解或误读不承担责任。\n"
]
}
],
"metadata": {
"colab": {
"collapsed_sections": [],
"name": "CBoW-PyTorch.ipynb",
"provenance": []
},
"interpreter": {
"hash": "16af2a8bbb083ea23e5e41c7f5787656b2ce26968575d8763f2c4b17f9cd711f"
},
"kernelspec": {
"display_name": "Python 3.8.12 ('py38')",
"language": "python",
"name": "python3"
},
"language_info": {
"codemirror_mode": {
"name": "ipython",
"version": 3
},
"file_extension": ".py",
"mimetype": "text/x-python",
"name": "python",
"nbconvert_exporter": "python",
"pygments_lexer": "ipython3",
"version": "3.8.12"
},
"orig_nbformat": 4,
"gpuClass": "standard",
"coopTranslator": {
"original_hash": "36df28efe3fe40b6fb0a7fa48fe3ea82",
"translation_date": "2025-08-31T10:36:21+00:00",
"source_file": "lessons/5-NLP/15-LanguageModeling/CBoW-PyTorch.ipynb",
"language_code": "zh"
}
},
"nbformat": 4,
"nbformat_minor": 0
}