You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

如何在PyTorch中使用预训练词嵌入?求快速实现方案

在PyTorch中快速使用词嵌入与预训练词嵌入的方法

一、从零训练自定义词嵌入

直接用PyTorch内置的nn.Embedding层就能快速实现,无需额外依赖:

  • 先构建词汇表(可通过统计语料词频生成,或用torchtext自动构建),给每个词分配唯一索引。
  • 初始化nn.Embedding层,传入词汇表大小和嵌入维度即可,训练时会自动反向更新嵌入参数。

代码示例:

import torch
import torch.nn as nn

vocab_size = 10000  # 假设词汇表包含10000个词
embed_dim = 128     # 设定嵌入维度为128

# 初始化嵌入层
embedding_layer = nn.Embedding(vocab_size, embed_dim)

# 示例输入:batch_size=2,每个样本含3个词索引
input_ids = torch.tensor([[1, 5, 3], [2, 4, 0]], dtype=torch.long)
# 获取词嵌入输出
embeddings = embedding_layer(input_ids)
print(embeddings.shape)  # 输出: torch.Size([2, 3, 128])

二、使用预训练词嵌入

最快的方式是借助PyTorch生态工具直接加载,或手动导入预训练文件:

方法1:用torchtext快速加载预训练嵌入(如GloVe)

torchtext支持一键加载常见预训练嵌入(GloVe、FastText等),自动匹配你的词汇表:

from torchtext.vocab import GloVe
import torch.nn as nn

# 加载GloVe 6B版本,嵌入维度100
glove_emb = GloVe(name='6B', dim=100)

# 假设已构建好词汇表vocab
vocab_size = len(vocab)
embed_dim = 100

# 初始化嵌入层并加载预训练权重
embedding_layer = nn.Embedding(vocab_size, embed_dim)
embedding_layer.weight.data.copy_(glove_emb.get_vecs_by_tokens(vocab.get_itos()))

# 若不想训练时更新预训练嵌入,冻结参数
embedding_layer.weight.requires_grad = False

方法2:手动加载预训练文件(如本地的Word2Vec/GloVe文件)

如果需要用自定义预训练文件,可按以下步骤实现:

import torch
import torch.nn as nn

# 读取预训练嵌入文件,构建词-向量字典
def load_pretrained_embeddings(embedding_path, vocab):
    emb_dict = {}
    with open(embedding_path, 'r', encoding='utf-8') as f:
        for line in f:
            values = line.strip().split()
            word = values[0]
            vec = torch.tensor([float(x) for x in values[1:]], dtype=torch.float)
            emb_dict[word] = vec

    # 初始化嵌入权重,未在预训练中出现的词用随机初始化
    embed_dim = len(next(iter(emb_dict.values())))
    embedding_weights = torch.randn(len(vocab), embed_dim)
    for idx, word in enumerate(vocab.get_itos()):
        if word in emb_dict:
            embedding_weights[idx] = emb_dict[word]
    return embedding_weights, embed_dim

# 加载本地预训练文件并初始化嵌入层
embedding_weights, embed_dim = load_pretrained_embeddings('glove.6B.100d.txt', vocab)
embedding_layer = nn.Embedding(len(vocab), embed_dim)
embedding_layer.weight.data.copy_(embedding_weights)
# 可选:冻结预训练参数
embedding_layer.weight.requires_grad = False

实用提示

  • 快速验证想法时,优先用nn.Embedding+torchtext预训练嵌入,避免手动处理文件的繁琐。
  • 训练自定义嵌入时,可设置padding_idx参数(如nn.Embedding(..., padding_idx=0)),让填充词的嵌入始终为0。
  • 常见预训练嵌入包括GloVe、Word2Vec、FastText,按需选择即可。

内容的提问来源于stack exchange,提问作者gotch4

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.07.23 08:12:11