如何创建适配不同尺寸数据的Sklearn风格神经网络类(含fit/predict)
适配任意输入维度的线性回归神经网络实现
核心问题分析
- 报错原因:原网络第一层输入维度硬编码为2,但你的训练/测试数据包含49个特征,导致矩阵乘法维度不匹配。
- 关于测试数据的担忧:只要训练和测试数据的特征列数一致(均为49),将网络输入层调整为对应维度后,测试数据完全可以正常预测,不存在适配问题。
修改后的可适配神经网络类
修改点:
- 允许初始化时指定输入特征维度
input_dim,动态构建输入层 - 修复梯度更新顺序错误(原代码先反向传播再清零梯度,会导致参数更新异常)
- 完善输入处理逻辑,自动兼容pandas DataFrame/Series和numpy数组输入
- 新增训练/评估模式切换,优化预测时的梯度计算逻辑
import torch import torch.nn as nn import torch.optim as optim import pandas as pd class MyNeuralNet(nn.Module): def __init__(self, input_dim, hidden_dim=4): super().__init__() # 动态适配输入特征维度 self.layer1 = nn.Linear(input_dim, hidden_dim, bias=True) self.layer2 = nn.Linear(hidden_dim, 1, bias=True) self.loss = nn.MSELoss() self.compile_() def forward(self, x): x = self.layer1(x) x = self.layer2(x) return x.squeeze() def fit(self, x, y, epochs=100): # 统一处理输入转tensor if isinstance(x, pd.DataFrame): x = torch.tensor(x.values, dtype=torch.float32) else: x = torch.tensor(x, dtype=torch.float32) if isinstance(y, (pd.DataFrame, pd.Series)): y = torch.tensor(y.values, dtype=torch.float32) else: y = torch.tensor(y, dtype=torch.float32) losses = [] for epoch in range(epochs): self.train() # 梯度更新正确流程:清零→前向→损失→反向→更新 self.opt.zero_grad() res = self.forward(x) loss_value = self.loss(res, y) loss_value.backward() self.opt.step() losses.append(loss_value.item()) return losses # 返回损失列表方便监控训练过程 def compile_(self, optimizer=optim.SGD, lr=0.01): # 支持自定义优化器和学习率 self.opt = optimizer(self.parameters(), lr=lr) def predict(self, x_test): self.eval() # 统一处理输入转tensor if isinstance(x_test, pd.DataFrame): x_test = torch.tensor(x_test.values, dtype=torch.float32) else: x_test = torch.tensor(x_test, dtype=torch.float32) with torch.no_grad(): # 评估阶段关闭梯度计算,提升效率 y_test_hat = self(x_test) return y_test_hat.detach().numpy()
使用示例
# 初始化模型时传入特征列数49 model = MyNeuralNet(input_dim=49) X_train = # pandas DataFrame,1168行49列 y_train = # pandas DataFrame/Series,1168行1列 X_test = # pandas DataFrame,292行49列 # 训练模型,可自定义训练轮数 losses = model.fit(X_train, y_train, epochs=150) # 生成预测结果 pred = model.predict(X_test) print(pred)
额外说明
- 网络的输入层只关注特征列数,与样本行数无关,因此1168行的训练数据和292行的测试数据都能正常处理(行数对应批量大小,不影响维度匹配)。
train()和eval()方法用于切换网络模式,若后续添加Dropout、BatchNorm等层时,能保证训练和评估时的行为正确。
内容的提问来源于stack exchange,提问作者leyjfk6
相关产品推荐
相关产品推荐

