如何修改PyTorch简易网络以适配不同维度输入数据
问题分析与解决方法
原代码存在的问题
- 语法与继承错误:
__init__方法里**input_dim的参数写法无效,应改为普通位置参数;同时super()中继承的类名写错,需和类定义一致为simpleNet。 - 维度不匹配核心原因:
nn.Linear层要求输入形状为[batch_size, feature_dim],若输入是高维数据(比如图片的[batch, channel, height, width]),直接传入会导致矩阵乘法维度不兼容——Linear会将输入除最后一维外的所有维度视为batch维度,若输入最后一维尺寸和input_dim不匹配,就会触发"mat1和mat2尺寸不匹配"错误。
修改后的代码方案
方案1:自动适配任意输入维度(无需提前指定input_dim)
import torch import torch.nn as nn class simpleNet(nn.Module): def __init__(self, hidden_size, num_classes): """ :param hidden_size: hidden dimension :param num_classes: total number of classes """ super(simpleNet, self).__init__() # 自动展平高维输入为[batch_size, 总特征数] self.flatten = nn.Flatten() # 延迟初始化hidden层,第一次forward时根据输入自动创建 self.hidden = None self.output = nn.Linear(hidden_size, num_classes) def forward(self, x): # 展平所有非batch维度到特征维度 x = self.flatten(x) # 第一次前向传播时,根据输入特征维度初始化hidden层 if self.hidden is None: input_dim = x.shape[1] self.hidden = nn.Linear(input_dim, hidden_size).to(x.device) # 执行前向计算 x = self.hidden(x) x = torch.sigmoid(x) x = self.output(x) return x
方案2:提前指定总特征数(更可控)
import torch import torch.nn as nn class simpleNet(nn.Module): def __init__(self, input_dim, hidden_size, num_classes): """ :param input_dim: 输入的总特征数(比如图片输入为C*H*W) :param hidden_size: hidden dimension :param num_classes: total number of classes """ super(simpleNet, self).__init__() self.flatten = nn.Flatten() self.hidden = nn.Linear(input_dim, hidden_size) self.output = nn.Linear(hidden_size, num_classes) def forward(self, x): # 展平输入,确保符合Linear层的输入要求 x = self.flatten(x) x = self.hidden(x) x = torch.sigmoid(x) x = self.output(x) return x
关键修改说明
- 添加
nn.Flatten()层:不管输入是2维表格数据([batch, feat])还是高维图片数据([batch, C, H, W]),都会自动将除第一个batch维度外的所有维度展平为一维特征,确保输入nn.Linear时形状合法。 - 延迟初始化(方案1):无需提前知道输入特征维度,第一次前向传播时自动根据输入创建适配的
hidden层,完全灵活处理任意输入维度。 - 修正基础错误:修复了类继承和参数定义的语法问题,避免不必要的报错。
使用示例
# 测试高维图片输入 model = simpleNet(hidden_size=128, num_classes=10) test_input = torch.randn(32, 3, 28, 28) # batch=32,3通道28x28图片 output = model(test_input) print(output.shape) # 输出torch.Size([32, 10]) # 测试2维表格输入 model2 = simpleNet(hidden_size=64, num_classes=2) test_input2 = torch.randn(16, 20) # batch=16,20维特征 output2 = model2(test_input2) print(output2.shape) # 输出torch.Size([16, 2])
内容的提问来源于stack exchange,提问作者Brenda Hu
相关产品推荐
相关产品推荐

