用简单神经网络拟合指数函数遇阻,求技术分析
问题背景
尝试训练神经网络拟合指数衰减函数 s(t) = s0 * e^(-t/decay_constant),网络输入为s和t,输出为s0与decay_constant。理论上神经网络应能较好逼近该函数,但实际效果远不如简单对数线性拟合。
实验设置
- 网络结构:带ReLU激活函数的稠密网络
- 损失函数:基于对数线性最小二乘法的MSE损失
- 训练数据:随机生成的指数衰减样本(CPU训练约20秒完成)
训练代码
import torch net = torch.nn.Sequential( torch.nn.Linear(10, 64), torch.nn.ReLU(), torch.nn.Linear(64, 64), torch.nn.ReLU(), torch.nn.Linear(64, 32), torch.nn.ReLU(), torch.nn.Linear(32, 2), torch.nn.ReLU(), # 期望输出始终为正 ) loss = torch.nn.MSELoss() optimizer = torch.optim.Adam(net.parameters(), lr=0.005) def signal_model(x, s0, decay): return s0 * torch.exp(-x / decay) # 生成数据点 batch_size = 4096 for episode in range(1000): # 随机生成时间、衰减常数、s0 t = (torch.arange(5, 55, 10.) + (torch.rand(batch_size, 5) * 10)).T decay_constant = torch.rand(batch_size) * 70 + 10 s0 = (torch.rand(batch_size) * 2 - 1) * 50 + 100 # 生成输入数据 y = signal_model(t, s0, decay_constant) data = torch.vstack((t, y)).T # 预测并计算损失 coefficients = net(data) l = loss(-torch.log(signal_model(t, *coefficients.T) + 1e-20), -torch.log(y + 1e-20)) optimizer.zero_grad(); l.backward(); optimizer.step()
可视化代码
import matplotlib.pyplot as plt t_ = torch.tensor([5, 15, 25, 35, 45], dtype=torch.float32) plotgrid = torch.arange(0, 50, 0.1) s = signal_model(t_, 20, 30) fig, ax = plt.subplots() ax.plot(plotgrid, signal_model(plotgrid, 20, 30).detach().numpy()) ax.scatter(t_, s.detach().numpy(), label="True") ax.scatter( t_, signal_model(t_, *net(torch.hstack((t_, s)).T)).detach().numpy(), label="Predicted", ) ax.legend()
已尝试的调整(均无改善)
- 调整学习率
- 修改网络层数/单元数量
- 调整批量大小
- 更换损失函数(系数MSE、无对数的信号MSE)
- 添加L1正则化
- 固定输入
t
原因分析
输出层激活函数错误:最后一层使用ReLU激活会导致严重问题。当模型输出接近0时,ReLU的梯度会变为0,参数无法更新;同时,衰减常数
decay_constant的取值范围是10-80,ReLU虽然能保证输出为正,但会引入不必要的非线性,而该问题本质是线性可解的(对数变换后),强行加非线性反而干扰拟合。损失函数与问题特性不匹配:原损失计算虽然用了对数变换,但结合ReLU的输出特性,梯度传播容易出现断层,无法有效引导模型收敛到最优解。此外,当前网络直接用10维输入预测全局参数,没有利用指数衰减函数的结构化线性特性。
网络结构冗余:该问题本质是对数线性可解的,无需复杂的深层ReLU网络,冗余的非线性层反而会增加拟合难度,导致模型陷入局部最优。
修正方案
1. 替换输出层激活函数
去掉最后一层的ReLU,改用Softplus(平滑版ReLU,保证输出为正且始终有梯度),或者直接对输出取指数来确保正性,避免梯度消失问题。
2. 利用问题的线性特性
对原函数取对数:log(s) = log(s0) - t/decay_constant,令a = log(s0),b = -1/decay_constant,则问题转化为线性回归log(s) = a + b*t。神经网络可以直接学习a和b,再通过逆变换得到s0 = exp(a),decay_constant = -1/b,大幅降低拟合难度。
3. 简化网络结构
由于问题本质是线性的,无需复杂的深层网络,使用1-2层线性层+少量隐藏单元即可满足需求,减少冗余非线性带来的干扰。
修正后的代码示例
import torch # 简化网络,输出层用Softplus保证正输出且梯度连续 net = torch.nn.Sequential( torch.nn.Linear(10, 32), torch.nn.ReLU(), torch.nn.Linear(32, 2), torch.nn.Softplus() ) loss = torch.nn.MSELoss() optimizer = torch.optim.Adam(net.parameters(), lr=0.001) # 降低学习率提升稳定性 def signal_model(x, s0, decay): return s0 * torch.exp(-x / decay) batch_size = 4096 for episode in range(1000): t = (torch.arange(5, 55, 10.) + (torch.rand(batch_size, 5) * 10)).T decay_constant = torch.rand(batch_size) * 70 + 10 s0 = (torch.rand(batch_size) * 2 - 1) * 50 + 100 y = signal_model(t, s0, decay_constant) data = torch.vstack((t, y)).T # 预测s0和decay,拆分维度匹配计算 pred_s0, pred_decay = net(data).T # 直接用对数域的MSE损失,贴合问题线性特性 log_y = torch.log(y + 1e-20) log_pred_y = torch.log(signal_model(t, pred_s0, pred_decay) + 1e-20) l = loss(log_pred_y, log_y) optimizer.zero_grad(); l.backward(); optimizer.step() # 每100轮打印损失监控训练状态 if episode % 100 == 0: print(f"Episode {episode}, Loss: {l.item():.4f}")
内容的提问来源于stack exchange,提问作者Daniel

