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

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.25 21:52:05