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

Word2Vec负采样代码报错求助,求协助修复训练功能

Word2Vec负采样代码报错修复求助

运行Word2Vec算法时触发如下报错:

RuntimeError: mat1 and mat2 shapes cannot be multiplied (2x4 and 100000x4)

问题卡在负采样代码部分,已知center_word和context_word为一维张量,需要修复代码使训练函数正常运行。

负采样函数代码

def negative_sampling(self, center_word, context_word):
        # 控制每个正样本对应的负样本数量
        num_neg_samples_per_center = self.num_neg_samples_per_center
        batch_size = center_word.shape[0]

        # 生成正样本对
        center_word_embeddings = self.center_embeddings(center_word) # batch_size, embedding_dim
        context_word_embeddings = self.context_embeddings(context_word) # batch_size, embedding_dim
        pos_pairs = torch.cat((center_word_embeddings.unsqueeze(1), context_word_embeddings.unsqueeze(1)), dim=1) # batch_size, 2, embedding_dim

        # 生成负样本
        neg_samples = []
        while len(neg_samples) < batch_size * num_neg_samples_per_center:
            # 基于词频随机采样
            neg_words = torch.multinomial(self.counts, num_neg_samples_per_center * batch_size, replacement=True)
            # 排除正样本词
            excluded = torch.cat((center_word.unsqueeze(1), context_word.unsqueeze(1)), dim=1)
            excluded = excluded.view(-1)
            neg_words = neg_words[~torch.isin(neg_words, excluded)]
            neg_samples += neg_words.tolist()
        neg_samples = neg_samples[:batch_size * num_neg_samples_per_center]
        neg_samples = torch.LongTensor(neg_samples).reshape(batch_size, num_neg_samples_per_center)

        # 获取负样本嵌入
        neg_samples_embeddings = self.context_embeddings(neg_samples) # batch_size, num_neg_samples_per_center, embedding_dim
        neg_samples_embeddings = neg_samples_embeddings.transpose(1,2) # batch_size, embedding_dim, num_neg_samples_per_center

        # 计算损失
        pos_scores = torch.bmm(pos_pairs, self.context_embeddings.weight.T.unsqueeze(0).expand(batch_size, -1, -1)).squeeze().sigmoid().log() # batch_size
        neg_scores = torch.bmm(neg_samples_embeddings.neg(), center_word_embeddings.unsqueeze(2)).squeeze().sigmoid().logsumexp(dim=1) # batch_size
        loss = -(pos_scores + neg_scores).mean()
        return loss

训练函数配置代码

run_training(
    model_type = 'neg', # 指定训练用的损失函数,'nll'为负对数损失,'neg'为负采样损失
    lr = 10, # 训练学习率
    num_neg_samples_per_center = 3, # 每个中心词对应的负样本数量
    checkpoint_model_path = './demo_checkpoints', # 模型 checkpoint 保存路径
    final_model_path = './final_demo_model', # 最终模型保存路径
    skip_window = 1, # 滑动窗口大小
    vocab_size = int(1e5), # 词汇表大小
    num_skips = 2, # 每个窗口采样的样本数
    batch_size = 256, # 训练批次大小(x,y对数量)
    embedding_size = 4, # 嵌入向量维度
    checkpoint_step = 500, # 每多少步保存一次 checkpoint
    max_num_steps = 2001 # 最大训练步数
)

问题分析与修复方案

核心报错原因

报错出现在正样本分数计算的torch.bmm步骤,张量维度不匹配:

  • pos_pairs维度为(batch_size, 2, embedding_dim)(例如256,2,4)
  • self.context_embeddings.weight.T.unsqueeze(0).expand(...)维度为(batch_size, embedding_dim, vocab_size)(例如256,4,100000)
  • bmm要求第一个张量的最后一维等于第二个张量的倒数第二维,但这里2≠4,导致矩阵乘法失败。

原代码的正样本分数计算逻辑完全错误,正确逻辑是计算单个中心词嵌入与对应上下文词嵌入的点积,而非和整个上下文嵌入矩阵做乘法。

修复后的完整负采样函数

def negative_sampling(self, center_word, context_word):
    num_neg_samples_per_center = self.num_neg_samples_per_center
    batch_size = center_word.shape[0]

    # 获取正样本嵌入
    center_emb = self.center_embeddings(center_word)  # (batch_size, embedding_dim)
    context_emb = self.context_embeddings(context_word)  # (batch_size, embedding_dim)

    # 优化负样本生成:减少循环次数,一次性采样后过滤正样本
    total_neg_needed = batch_size * num_neg_samples_per_center
    neg_samples = []
    while len(neg_samples) < total_neg_needed:
        # 多采样50%的候选样本,避免多次循环
        neg_candidates = torch.multinomial(self.counts, int(total_neg_needed * 1.5), replacement=True)
        # 合并需要排除的正样本词(中心词+上下文词)
        excluded_words = torch.cat([center_word, context_word])
        # 过滤掉正样本词
        valid_neg = neg_candidates[~torch.isin(neg_candidates, excluded_words)]
        neg_samples.extend(valid_neg.tolist())
    # 截取所需数量并调整维度
    neg_samples = torch.LongTensor(neg_samples[:total_neg_needed]).reshape(batch_size, num_neg_samples_per_center)

    # 获取负样本嵌入
    neg_emb = self.context_embeddings(neg_samples)  # (batch_size, num_neg, embedding_dim)
    neg_emb = neg_emb.transpose(1, 2)  # (batch_size, embedding_dim, num_neg)

    # 计算正样本分数:中心词与对应上下文词的点积
    pos_scores = torch.sum(center_emb * context_emb, dim=1).sigmoid().log()  # (batch_size,)

    # 计算负样本分数:中心词与所有负样本的点积取负后计算logsumexp
    neg_scores = torch.bmm(-neg_emb, center_emb.unsqueeze(2)).squeeze().sigmoid().logsumexp(dim=1)  # (batch_size,)

    # 总损失
    loss = -(pos_scores + neg_scores).mean()
    return loss

内容的提问来源于stack exchange,提问作者Blue And Red

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.29 09:35:35