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

如何在PyTorch中实现不同维度的嵌入张量与掩码张量相乘?

问题解决

先修复代码中的明显错误

  1. 模块初始化缺失:__init__方法必须调用父类初始化函数,否则PyTorch无法正确注册模块参数,添加super().__init__():
class Network(nn.Module):
    def __init__(self, embedding_dim, hidden_dim, n_outputs):
        super().__init__()  # 必须添加这一行
        self.embedding = nn.Embedding(256, embedding_dim)
        self.linear1 = nn.Linear(embedding_dim, hidden_dim)
        self.linear2 = nn.Linear(hidden_dim, n_outputs)
  1. 变量名拼写错误:forward方法最后一行把drop_out误写为dropout,修正后:
out = self.linear2(drop_out)

维度不匹配问题排查

你遇到的RuntimeError提示第3维度(索引从0开始)尺寸16和22不匹配,说明操作涉及的两个张量存在4维维度冲突,核心问题大概率是输入张量维度不符合预期,按以下步骤检查:

  • 确认输入chars的形状:chars必须是[8,22]的整数张量(批量大小8,每个单词22个字符),这样经过nn.Embedding后才会得到[8,22,16]的嵌入张量。如果chars是4维张量(比如[8, x, y, 22]),嵌入后会变成[8,x,y,22,16],和mask运算时就会触发维度不匹配错误。

  • 确认输入mask的形状:mask必须是[8,22],unsqueeze(2)后变成[8,22,1],PyTorch的广播机制会自动把它扩展到[8,22,16],和字符嵌入张量逐元素相乘是合法的。如果mask形状不对(比如[8,16]),也会导致维度冲突。

修正后的完整代码

import torch
import torch.nn as nn
import torch.nn.functional as F

class Network(nn.Module):
    def __init__(self, embedding_dim, hidden_dim, n_outputs):
        super().__init__()
        self.embedding = nn.Embedding(256, embedding_dim)
        self.linear1 = nn.Linear(embedding_dim, hidden_dim)
        self.linear2 = nn.Linear(hidden_dim, n_outputs)

    def forward(self, chars, mask):
        # chars形状: [8,22], mask形状: [8,22]
        chars_embeddings = self.embedding(chars)  # 输出形状: [8,22,embedding_dim]
        # mask扩展为[8,22,1],和chars_embeddings广播相乘
        chars_embeddings = torch.mul(chars_embeddings, mask.unsqueeze(2))
        # 对字符维度做均值池化,得到每个单词的嵌入
        pool_embeddings = chars_embeddings.mean(dim=1)  # 输出形状: [8,embedding_dim]
        lin_out = self.linear1(pool_embeddings)
        drop_out = F.dropout(F.relu(lin_out), training=self.training)
        out = self.linear2(drop_out)
        return out

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.11 13:17:12