如何在PyTorch中实现不同维度的嵌入张量与掩码张量相乘?
问题解决
先修复代码中的明显错误
- 模块初始化缺失:
__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)
- 变量名拼写错误: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
相关产品推荐
相关产品推荐

