使用JAX最小化两点Lennard-Jones势及力时结果不符问题排查
问题分析
你遇到的核心问题是:最小化Lennard-Jones势的力平方时,BFGS优化器收敛到了r→∞的区域,而非预期的平衡位置r=√2≈1.41。
原因拆解
Lennard-Jones势的力平方函数(dU/dr)²存在两个极小值区域:
- 平衡位置
r=√2:此时dU/dr=0,力平方为0,是局部极小值。 r→∞:此时dU/dr趋近于0,力平方也趋近于0,是全局极小值。
从初始点x=[0,1](对应r=1)出发,BFGS算法会沿着函数值下降最快的方向迭代:当r>√2时,力平方函数随r增大持续递减(趋近于0),且梯度绝对值越来越小,算法会自然收敛到r极大的区域,最终得到你看到的r=10的结果。
修正方案
方案1:调整初始值
将初始位置设置为更接近平衡位置的值,引导优化器收敛到预期的局部极小值:
x_init = jnp.array([0.0, 1.4], dtype=jnp.float64) # 接近r=√2≈1.414
方案2:修改目标函数,消除全局极小值
在力平方的目标函数中加入惩罚项,抑制r过大的情况。例如加入1/r项,让r→∞时函数值不再趋近于0:
def force(r): d_potential_dr = grad(potential) return (d_potential_dr(r)**2) + 1/r # 加入惩罚项
方案3:优化势能而非力平方
势能的全局极小值唯一对应平衡位置,直接优化势能是更可靠的方式(你已经验证过这种方法有效)。
代码优化小细节
每次调用force(r)时重新计算grad(potential)会造成冗余,建议提前预计算梯度函数:
d_potential_dr = grad(potential) # 提前计算一次梯度函数 def force(r): r = jnp.where(r == 0, jnp.finfo(jnp.float64).eps, r) # 避免r=0的情况 return (d_potential_dr(r)**2)
内容的提问来源于stack exchange,提问作者Heng Yuan
相关产品推荐
相关产品推荐

