PyTorch输入维度不匹配却能前向传播的问题排查
PyTorch低版本Linear层维度不匹配未报错问题
搭建自编码器模型时发现异常:输入维度与Linear层定义的输入维度不匹配时,PyTorch 1.9.0+cu111版本中前向传播居然能正常执行,而高版本会触发预期的维度不匹配错误。
模型定义
import torch import torch.nn as nn device = torch.device("cuda" if torch.cuda.is_available() else "cpu") class Encoder(nn.Module): def __init__(self,input_dim,latent_dim): super().__init__() self.linear1 = nn.Linear(input_dim,64) self.linear2 = nn.Linear(64,16) self.linear3 = nn.Linear(16,latent_dim) self.relu = nn.ReLU() def forward(self,x): out = self.linear1(x) out = self.relu(out) out = self.linear2(out) out = self.relu(out) latent = self.linear3(out) return latent class Decoder(nn.Module): def __init__(self, latent_dim, output_dim): super().__init__() self.linear1 = nn.Linear(latent_dim,16) self.linear2 = nn.Linear(16,64) self.linear3 = nn.Linear(64,output_dim) self.relu = nn.ReLU() def forward(self,x): out = self.linear1(x) out = self.relu(out) out = self.linear2(out) out = self.relu(out) result = self.linear3(out) return result class AutoEncoder(nn.Module): def __init__(self,input_dim, latent_dim, output_dim): super().__init__() self.encoder = Encoder(input_dim,latent_dim) self.decoder = Decoder(latent_dim,output_dim) def forward(self,x): enc = self.encoder(x) dec = self.decoder(enc) return dec model_AE = AutoEncoder(input_dim = 13, latent_dim = 8, output_dim = 13).to(device)
测试情况
- Test1:输入形状为
torch.Size([254, 13]),与模型input_dim=13匹配,前向传播正常,符合预期。 - Test2:输入由两个13维张量拼接而成,形状为
torch.Size([254, 26]),与第一层Linear的输入维度13不匹配,但PyTorch 1.9.0+cu111中前向传播正常执行,不符合预期。
跨版本测试简化代码
import torch import torch.nn as nn device = torch.device("cuda" if torch.cuda.is_available() else "cpu") class Encoder(nn.Module): def __init__(self,input_dim,latent_dim): super().__init__() self.linear1 = nn.Linear(input_dim,64) self.linear2 = nn.Linear(64,16) self.linear3 = nn.Linear(16,latent_dim) self.relu = nn.ReLU() def forward(self,x): out = self.linear1(x) out = self.relu(out) out = self.linear2(out) out = self.relu(out) latent = self.linear3(out) return latent class Decoder(nn.Module): def __init__(self, latent_dim, output_dim): super().__init__() self.linear1 = nn.Linear(latent_dim,16) self.linear2 = nn.Linear(16,64) self.linear3 = nn.Linear(64,output_dim) self.relu = nn.ReLU() def forward(self,x): out = self.linear1(x) out = self.relu(out) out = self.linear2(out) out = self.relu(out) result = self.linear3(out) return result class AutoEncoder(nn.Module): def __init__(self,input_dim, latent_dim, output_dim): super().__init__() self.encoder = Encoder(input_dim,latent_dim) self.decoder = Decoder(latent_dim,output_dim) def forward(self,x): enc = self.encoder(x) dec = self.decoder(enc) return dec model_AE = AutoEncoder(input_dim = 13, latent_dim = 8, output_dim = 13).to(device) inputxx = torch.randn(254, 26).to(device) out = model_AE(inputxx)
测试结果
- torch版本1.9.0+cu111:代码正常执行(不符合预期)
- torch版本2.1.2+cu121:代码触发维度不匹配错误(符合预期)
结论:该现象是PyTorch 1.9.0+cu111版本存在的bug。
内容的提问来源于stack exchange,提问作者bubba
相关产品推荐
相关产品推荐

