基于 PyTorch 实现 CBOW 词向量模型

一、什么是 CBOW

CBOW(Continuous Bag of Words)是 Word2Vec 中的一种经典模型,主要用于学习单词的词向量。

它的基本思想是:

根据一个单词周围的上下文单词,预测中间的目标单词。

例如:

复制代码
People create programs to direct processes

如果设置上下文窗口大小为 2,那么模型可以利用目标词前后各两个单词来预测中间的单词。

CBOW 的训练过程实际上就是让神经网络不断学习"上下文单词"和"目标单词"之间的关系,最终得到每个单词对应的低维向量。


二、导入相关库

代码首先导入 PyTorch、NumPy 和 tqdm:

复制代码
import torch
import torch.optim as optim
import torch.nn as nn
import torch.nn.functional as F
from tqdm import tqdm
import numpy as np

其中:

  • torch:PyTorch 深度学习框架。

  • torch.nn:用于构建神经网络。

  • torch.optim:提供 Adam 等优化器。

  • torch.nn.functional:提供 relulog_softmax 等函数。

  • tqdm:显示训练进度。

  • numpy:用于词向量的保存和处理。


三、准备语料库

代码使用一段英文文本作为训练语料:

复制代码
raw_text = """We are about to study the idea of a computational process.
Computational processes are abstract beings that inhabit computers.
As they evolve, processes manipulate other abstract things called data.
The evolution of a process is directed by a pattern of rules
called a program. People create programs to direct processes. In effect,
we conjure the spirits of the computer with our spells.""".split()

使用 split() 可以按照空格将文本划分成一个个单词。

例如:

复制代码
We are about to study

会变成:

复制代码
['We', 'are', 'about', 'to', 'study']

然后设置上下文窗口:

复制代码
CONTEXT_SIZE = 2

这里的 2 表示目标单词左右各取两个单词作为上下文。


四、建立词表

首先使用 set() 获取语料库中的不重复单词:

复制代码
vocab = set(raw_text)
vocab_size = len(vocab)

vocab 就是整个语料库的词汇表,而 vocab_size 表示词汇表大小。

接着建立两个字典:

复制代码
word_to_idx = {word:i for i,word in enumerate(vocab)}
idx_to_word = {i:word for i,word in enumerate(vocab)}

其中:

复制代码
word_to_idx

负责把单词转换成数字编号。

例如:

复制代码
People → 5
create → 10
programs → 8

而:

复制代码
idx_to_word

负责把数字编号转换回单词。

这种"单词 ↔ 数字"的转换非常重要,因为神经网络不能直接处理字符串,需要先把单词转换成数字。


五、构造 CBOW 训练数据

CBOW 最重要的一步就是构造训练数据。

代码:

复制代码
for i in range(CONTEXT_SIZE,len(raw_text)-CONTEXT_SIZE):

不断选择一个单词作为目标词,然后取它前后的单词作为上下文。

例如:

复制代码
People create programs to direct

如果目标词是:

复制代码
programs

那么上下文可以是:

复制代码
People create to direct

代码最终生成:

复制代码
(context, target)

这样的训练样本。

例如:

复制代码
输入:
['People', 'create', 'to', 'direct']

目标:
programs

因此 CBOW 的训练任务可以简单理解为:

复制代码
上下文单词
   ↓
词向量
   ↓
平均
   ↓
神经网络
   ↓
预测目标单词

代码中的训练数据构造过程就是围绕这个思想完成的。


六、将单词转换成 Tensor

构造好训练数据后,还需要把单词转换成数字:

复制代码
def make_context_vector(context,word_to_idx):
    idxs = [word_to_idx[x] for x in context]
    return torch.tensor(idxs,dtype=torch.long)

例如:

复制代码
['People', 'create', 'to', 'direct']

经过 word_to_idx 后可能变成:

复制代码
[3, 7, 12, 5]

然后转换成:

复制代码
tensor([3, 7, 12, 5])

这里使用 torch.long 是因为 nn.Embedding 的输入需要使用整数索引。


七、选择 CPU 还是 GPU

代码使用:

复制代码
device = 'cuda' if torch.cuda.is_available() else \
         'mps' if torch.backends.mps.is_available() else 'cpu'

自动判断当前电脑是否支持 GPU。

如果有 NVIDIA CUDA:

复制代码
cuda

如果没有,则继续判断 Apple 的 MPS,最后使用:

复制代码
cpu

模型也会被放到对应设备:

复制代码
model = CBOW(vocab_size,10).to(device)

训练数据也需要放到相同设备:

复制代码
context_vector = make_context_vector(
    context,word_to_idx
).to(device)

这样可以避免模型和数据分别位于 CPU、GPU 而导致的设备错误。


八、CBOW 神经网络结构

代码定义了一个 CBOW 类:

复制代码
class CBOW(nn.Module):

继承 nn.Module 是 PyTorch 构建神经网络的基本方式。

模型主要包含三层:

复制代码
self.embedding = nn.Embedding(vocab_size, embedding_dim)

self.proj = nn.Linear(embedding_dim,128)

self.output = nn.Linear(128,vocab_size)

整体结构可以表示为:

复制代码
输入单词编号
      ↓
Embedding
      ↓
多个词向量
      ↓
求平均
      ↓
Linear
      ↓
ReLU
      ↓
Linear
      ↓
词汇表大小
      ↓
预测目标单词

九、Embedding 词嵌入

代码:

复制代码
self.embedding = nn.Embedding(vocab_size,embedding_dim)

这里的:

复制代码
embedding_dim = 10

表示每个单词最终使用一个长度为 10 的向量表示。

例如:

复制代码
People → [0.12, -0.35, ..., 0.27]

原本一个单词只是一个编号,而经过 Embedding

相关推荐
JarmanYuo2 小时前
YOLO 涨点研究(十二):具身 CV 进阶篇——Sim2Real 域随机化与真机部署
人工智能·pytorch·python·yolo·计算机视觉
马剑威(威哥爱编程)15 小时前
【共创稿事节】HarmonyOS 7 应用 Skill 化实战:从“被打开“到“被调用“,把功能递进系统意图分发池
pytorch·深度学习·harmonyos
魔镜er1 天前
11-循环神经网络
人工智能·pytorch·python·深度学习·神经网络
派大_星1 天前
卷积神经网络(CNN)学习笔记:从卷积原理到ResNet与迁移学习
pytorch
zx_741484811 天前
【深度学习入门】PyTorch 食物图像分类三连:CNN 训练、学习率调度与 ResNet18 迁移学习
pytorch·深度学习·cnn·迁移学习
Zguigo2 天前
【CUDA1】GPUvsCPU,CUDA Kernel
人工智能·pytorch·深度学习
Jared_devin2 天前
单头、多头 Self-Attention可视化
人工智能·pytorch·深度学习·算法·机器学习·pycharm
昇腾知识体系2 天前
昇腾环境安装 flash_attn:FlashAttnPrefillBackend 报错与替代算子
人工智能·pytorch·华为·知识图谱
维基框架3 天前
PyTorch正在重构开源AI基础设施
人工智能·pytorch·重构