PyTorch中MaxPool1d后矩阵无法相乘错误的解决求助
解决建议
你的错误是全连接层输入维度的手动计算值与实际经过卷积池化后的特征维度不匹配,导致矩阵乘法无法执行。以下是具体的排查和解决方法:
1. 先确认实际特征维度
在forward方法中添加打印语句,输出每一层的张量形状,直观看到维度变化:
def forward(self, xb: torch.Tensor): xb = xb.permute(0, 2, 1) print("Permuted input shape:", xb.shape) current = xb for name, module in self.conv_net.named_children(): current = module(current) print(f"After {name}: {current.shape}") return current
运行后观察flatten层的输出形状,比如如果输出是(batch_size, 2048),全连接层的in_features就应设为2048,而非手动计算的2000。
2. 动态创建全连接层(推荐)
避免手动计算维度误差,第一次前向传播时根据实际特征维度动态创建全连接层:
class DNA_CNN_test2(nn.Module): def __init__(self, seq_len: int, num_filters: List[int] = [32, 64,128], kernel_size: int = 3, p = 0.2): super().__init__() self.seq_len = seq_len # CNN module self.conv_net = nn.Sequential() num_filters = [4] + num_filters for idx in range(len(num_filters) - 1): self.conv_net.add_module( f"conv_{idx}", nn.Conv1d(num_filters[idx], num_filters[idx + 1], kernel_size=kernel_size, padding='same') ) self.conv_net.add_module(f"relu_{idx}", nn.ReLU(inplace=True)) self.conv_net.add_module(f"batchNor_{idx}",nn.BatchNorm1d(num_filters[idx + 1])) self.conv_net.add_module(f"dropout_{idx}",nn.Dropout(0.2)) self.conv_net.add_module(f"MaxP_{idx}",nn.MaxPool1d(4,stride= 4)) self.conv_net.add_module("flatten", nn.Flatten()) # 先不初始化全连接层 self.linear = None def forward(self, xb: torch.Tensor): """Forward pass.""" xb = xb.permute(0, 2, 1) out = self.conv_net(xb) if self.linear is None: # 根据实际特征维度创建全连接层 self.linear = nn.Linear(out.shape[1], 1).to(xb.device) out = self.linear(out) return out
这种方式自动适配输入序列长度变化,无需手动计算池化后维度。
3. 对齐Keras与PyTorch的池化行为
如果原Keras代码中的MaxPooling1D使用了padding='same',PyTorch默认的MaxPool1d(padding=0)会导致输出维度不一致。此时需给PyTorch的MaxPool1d添加padding参数:
# 对应Keras的MaxPooling1D(pool_size=4, strides=4, padding='same') nn.MaxPool1d(4, stride=4, padding='same')
这样池化后的输出长度与Keras保持一致,手动计算的维度也会更准确。
4. 使用自适应池化简化维度计算
如果不需要严格对齐原Keras的池化行为,可改用AdaptiveMaxPool1d固定特征长度为1,全连接层输入维度直接取最后一层滤波器数量:
class DNA_CNN_test2(nn.Module): def __init__(self, seq_len: int, num_filters: List[int] = [32, 64,128], kernel_size: int = 3, p = 0.2): super().__init__() self.seq_len = seq_len # CNN module self.conv_net = nn.Sequential() num_filters = [4] + num_filters for idx in range(len(num_filters) - 1): self.conv_net.add_module( f"conv_{idx}", nn.Conv1d(num_filters[idx], num_filters[idx + 1], kernel_size=kernel_size, padding='same') ) self.conv_net.add_module(f"relu_{idx}", nn.ReLU(inplace=True)) self.conv_net.add_module(f"batchNor_{idx}",nn.BatchNorm1d(num_filters[idx + 1])) self.conv_net.add_module(f"dropout_{idx}",nn.Dropout(0.2)) # 用自适应池化替代多次MaxPool1d,固定输出长度为1 self.conv_net.add_module("adaptive_pool", nn.AdaptiveMaxPool1d(1)) self.conv_net.add_module("flatten", nn.Flatten()) # 全连接层输入维度为最后一层滤波器数量 self.conv_net.add_module("linear",nn.Linear(num_filters[-1], 1)) def forward(self, xb: torch.Tensor): """Forward pass.""" xb = xb.permute(0, 2, 1) out = self.conv_net(xb) return out
内容的提问来源于stack exchange,提问作者Jin_soo
相关产品推荐
相关产品推荐

