基于JAX的机器学习模型完美拟合优化方法咨询
阻尼振荡拟合模型优化与损失函数选择建议
我是机器学习新手,已经用JAX实现了一个阻尼振荡拟合模型,目前的拟合效果有附图参考,想知道怎么进一步优化模型实现完美拟合;另外我在高效选择损失函数方面没什么经验。我的代码是在已有函数拟合程序基础上,结合带噪声的电压测量数据文本文件编写的,代码如下:
#!/usr/bin/env python3 # # Fit function to data import matplotlib.pyplot as plt import numpy as np import jax.numpy as jnp from jax import grad, jit, vmap, random # load some noisy data test = np.loadtxt('newe.txt') N = 200 sigma = 0.05 x = test[:, 0] y = test[:, 1] #plt.plot(x,y) #plt.show() # Match function to data def func(params, x): # Parameterised damped oscillation l, omega = params # Note, we "normalise" parameters y_pred = jnp.exp(l*10 * x) * jnp.sin(2*jnp.pi* omega*10 * x) return y_pred def loss(params, x, y): # Loss function y_pred = func(params, x) return jnp.mean((y - y_pred)**2) # Compile loss and gradient c_loss = jit(loss) d_loss = jit(grad(loss)) # One iteration of gradient descent def update_params(params, x, y): grads = d_loss(params, x, y) params = [param - 0.1 * grad for param, grad in zip (params, grads)] return params # Initialise parameters key = random.PRNGKey(0) params = [random.normal(key, (1,)), random.normal(key, (1,))] err = [] for epoch in range(100000): err.append(c_loss(params, x, y)) params = update_params(params, x, y) err.append(c_loss(params, x, y)) print("Damping: ", params[0]*10) print("Frequency:", params[1]*10) y_pred = func(params, x) # Plot loss and predictions f, ax = plt.subplots(1,2) ax[0].semilogy(err) ax[0].set_title("History") ax[1].plot(x, y, label="ground truth") ax[1].plot(x, y_pred, label="predictions") ax[1].legend() plt.show()
一、模型优化方案
- 优化参数初始化:当前用随机正态分布初始化参数,容易陷入局部最优。可以先可视化数据手动估算初始值:从振荡周期算出频率初始值,从衰减幅度估算阻尼系数初始值,让参数起点更接近真实值,加快收敛并避免局部最优。
- 替换自适应优化器:固定步长梯度下降易出现收敛慢或震荡问题,建议用JAX生态的
optax库中的自适应优化器(如Adam、带动量的SGD),这类优化器能自动调整学习率,收敛更稳定高效。示例替换代码:import optax # 定义优化器 optimizer = optax.adam(learning_rate=0.01) opt_state = optimizer.init(params) # 更新函数替换 def update_params(params, opt_state, x, y): grads = d_loss(params, x, y) updates, opt_state = optimizer.update(grads, opt_state) params = optax.apply_updates(params, updates) return params, opt_state - 调整参数归一化逻辑:当前将参数乘10的“归一化”会干扰梯度尺度,建议直接将参数尺度融入模型公式(如
y_pred = jnp.exp(l * x) * jnp.sin(2*jnp.pi*omega * x)),再配合优化器调整学习率,让参数物理意义更明确,梯度计算更合理。 - 添加训练终止条件:固定10万轮训练会浪费资源,可设置终止阈值,比如当连续20轮损失下降小于
1e-8时停止训练,避免无效迭代。 - 完善模型结构:若数据存在直流偏移或相位偏移,需扩展模型为
y_pred = C + jnp.exp(l * x) * jnp.sin(2*jnp.pi*omega * x + phi),加入常数项C和相位phi,提升拟合能力。
二、损失函数选择经验
- 均方误差(MSE):你当前使用的MSE适合高斯噪声场景(多数电压测量噪声符合高斯分布),对大误差惩罚更重,但数据存在异常值时会被严重干扰。
- 平均绝对误差(MAE):若数据中有较多异常点,MAE鲁棒性更强,对误差的惩罚是线性的,不会被极端值带偏。JAX实现:
jnp.mean(jnp.abs(y - y_pred))。 - Huber损失:结合MSE和MAE的优势,小误差时用MSE(二阶可导,梯度稳定),大误差时切换为MAE(抗异常值),适合混合噪声场景。自定义实现:
def huber_loss(params, x, y, delta=1.0): y_pred = func(params, x) residual = jnp.abs(y - y_pred) return jnp.mean(jnp.where(residual < delta, 0.5 * residual**2, delta * (residual - 0.5 * delta))) - 选择原则:
- 优先判断噪声类型:高斯噪声选MSE,拉普拉斯噪声/含异常值选MAE或Huber。
- 看拟合目标:追求预测精度且无异常值用MSE;追求抗干扰能力用MAE/Huber。
- 考虑可导性:MSE和Huber小误差区域二阶可导,梯度稳定,适配梯度下降类优化器;MAE一阶可导,零点梯度不连续,需搭配支持的优化器。
内容的提问来源于stack exchange,提问作者jeevan
相关产品推荐
相关产品推荐

