PyTorch特征拼接后矩阵维度不匹配RuntimeError问题求解
问题解决方案
核心问题分析
报错的本质是线性层输入维度不匹配:拼接后的特征张量形状是[10,2048],但fin_old线性层定义为nn.Linear(64,2),要求输入维度为64,矩阵乘法时2048≠64,导致形状不兼容。同时需确保所有待拼接张量的设备(CPU/GPU)一致,避免隐性错误。
具体修复步骤
- 修正线性层输入维度:将
fin_old的输入维度改为拼接后的总特征数(768+512+768=2048),即定义为nn.Linear(2048, 2)。若存在其他同类型线性层,需同步修改输入维度与拼接后的特征维度一致。 - 统一张量设备:拼接前确保
x、y、rag在同一设备上,可通过.to(device)强制对齐:device = x.device y = y.to(device) rag = rag.to(device) - 确认拼接维度:使用
torch.cat时指定dim=1(按特征维度拼接),避免因维度参数错误导致形状异常。
示例代码修改
原错误代码:
class Classifier(nn.Module): def __init__(self): super().__init__() self.fin_old = nn.Linear(64, 2) # 输入维度不匹配 def forward(self, x, y, rag): concatenated = torch.cat([x, y, rag], dim=1) output = self.fin_old(concatenated) # 触发维度不匹配报错 return output
修改后代码:
class Classifier(nn.Module): def __init__(self): super().__init__() self.fin_old = nn.Linear(2048, 2) # 匹配拼接后的特征维度 def forward(self, x, y, rag): # 统一设备 device = x.device y = y.to(device) rag = rag.to(device) # 按特征维度拼接 concatenated = torch.cat([x, y, rag], dim=1) output = self.fin_old(concatenated) return output
内容的提问来源于stack exchange,提问作者Harshal Dharpure
相关产品推荐
相关产品推荐

