PyTorch回归模型无法拟合y=x²,求故障排查方案
模型无法拟合y=x²曲线的问题排查
我在使用PyTorch开发机器学习回归项目时,遇到模型无法学习的问题:模型始终输出近乎直线,损失几乎没有下降。为定位问题,我将原项目简化为拟合y=x²曲线的最小复现程序,但问题依旧。
该简化程序的逻辑如下:
- 复用原项目中支持6个特征的ANN模型类
- 训练时前5个特征固定为0,最后一个特征取区间[-2,2]内的均匀分布数值x
- 通过
generate_targets()生成对应的目标值y=x²
复现代码
from torch import tensor, float32 from torch import nn import torch.optim as optim from torch.utils.data import Dataset, DataLoader import matplotlib.pyplot as plt class ANN(nn.Module): def __init__(self, feature_num: int): super(ANN,self).__init__() self.layers = nn.Sequential( nn.Linear(feature_num, 300), nn.Tanh(), nn.Linear(300, 200), nn.Tanh(), nn.Linear(200, 150), nn.Tanh(), nn.Linear(150, 50), nn.Tanh(), nn.Linear(50, 1) ) def forward(self, x): predictions = self.layers(x) return predictions class TestDataset(Dataset): def __init__(self, sample_num): self.sample_num = sample_num self.func_max = 2 self.func_min = -2 self.unit = (self.func_max - self.func_min) / self.sample_num self.targets = generate_targets(self.sample_num) def __getitem__(self, index): x = self.func_min + index * self.unit return tensor([0, 0, 0, 0, 0, x], dtype=float32), self.targets[index] def __len__(self): return self.sample_num # Generate the list of y def generate_targets(count): func_max = 2 func_min = -2 unit = (func_max - func_min)/count target_list = [] for i in range(count): x = func_min + unit*i y = x ** 2 target_list.append(y) return target_list # The main program def start_train(): sample_num = 500 train_data = TestDataset(sample_num) train_dataloader = DataLoader(train_data, batch_size=10, shuffle=True) model = ANN(6) model.train() mae_loss = nn.L1Loss() optimizer = optim.Adam(model.parameters()) loss_list = [] for i in range(sample_num): train_feature, train_target = next(iter(train_dataloader)) prediction = model(train_feature) loss = mae_loss(prediction, train_target.float().unsqueeze(1)) loss.backward() optimizer.step() optimizer.zero_grad() loss_list += [loss.item()] print(f"iteration [{i + 1}/{sample_num}] Loss = {loss.item():.3f}") plt.plot(range(sample_num), loss_list, marker='o', label='Validation') plt.xlabel('iterations') plt.ylabel('MAE loss') plt.title('Loss vs Iteration') plt.legend(loc='upper right') plt.savefig('Debugger_Loss_Iter.png') plt.close() comparison(model, sample_num) # Evaluate the trained model by plugging in each x coord and see the generated comparative graph def comparison(model: ANN, sample_num: int) -> None: model.eval() prediction_list = [] for i in range(sample_num): train_feature = tensor([0, 0, 0, 0, 0, i / sample_num], dtype=float32) prediction = model(train_feature) prediction_list += [prediction.item()] target_list = generate_targets(sample_num) x_list = [(i - sample_num / 2) / sample_num for i in list(range(sample_num))] plt.plot(x_list, target_list, marker='o', label='Target') plt.plot(x_list, prediction_list, marker='o', label='Prediction') plt.xlabel('x') plt.ylabel('y') plt.title('Prediction-target Comparison') plt.legend(loc='upper right') plt.savefig('Debugger_Comparison.png') plt.close() if __name__ == '__main__': start_train()
当前现象
- 模型输出近乎直线,与目标y=x²曲线差距极大
- 损失曲线几乎没有下降趋势
已排查内容
已确认所有张量的形状和数据类型均符合预期:
train: train_feature: (10,6) float32 train_target.shape([10,1]) float32 prediction: ([10,1]) float32 comparison: train_feature:([10,6]) float32 prediction([10,1]) float32
恳请帮忙排查模型无法学习的原因。
内容的提问来源于stack exchange,提问作者G2bbJ D9jd5
相关产品推荐
相关产品推荐

