使用JAX优化两点Lennard-Jones势能的Python代码结果异常
问题诊断与解决方案
核心问题:平动不变性导致优化歧义
你的Lennard-Jones势能仅与两点间的相对距离有关,与整体位置无关——也就是说将两个点同时平移任意距离,势能值完全不变。BFGS这类无约束优化器会利用这个平动自由度,在优化过程中调整整体位置,而非聚焦于寻找势能最小的最优距离。初始位置[0,1]对应的势能为0,而最优距离≈1.12时势能为-2(远低于0),但优化器因平动自由度的干扰,误将势能平缓趋近于0的大距离区域当成了最小值点。
解决方案:消除平动自由度
通过固定位置或添加约束消除平动自由度,让优化器仅关注两点间的相对距离,即可得到正确结果。
方法1:固定单个点的位置
直接固定其中一个点的坐标,只优化另一个点的位置,彻底消除平动自由度:
import jax import jax.numpy as jnp from jax.scipy.optimize import minimize jax.config.update("jax_enable_x64", True) # 仅优化第二个点的位置,初始值为1 x_init = jnp.array([1.0], dtype=jnp.float64) epsilon = 1 sigma = 1 def potential(r): r = jnp.where(r == 0, jnp.finfo(jnp.float64).eps, r) return 4 * epsilon * ((sigma/r)**12 - (sigma/r)**6) def F(x): # 第一个点固定在0,计算与第二个点的距离 r = jnp.abs(x[0] - 0) return potential(r) result = minimize(F, x_init, method='BFGS') print("优化后的第二个点位置:", result.x[0]) print("两点距离:", result.x[0]) print("理论最优距离:", 2**(1/6))
方法2:添加等式约束(使用SLSQP方法)
保留两个点的优化,但添加约束固定整体位置(比如两点坐标和与初始值一致),用支持约束的SLSQP方法求解:
import jax import jax.numpy as jnp from jax.scipy.optimize import minimize N = 2 jax.config.update("jax_enable_x64", True) x_init = jnp.arange(N, dtype=jnp.float64) epsilon = 1 sigma = 1 def potential(r): r = jnp.where(r == 0, jnp.finfo(jnp.float64).eps, r) return 4 * epsilon * ((sigma/r)**12 - (sigma/r)**6) def F(x): r = jnp.abs(x[:, None] - x[None, :]) pot = jax.vmap(jax.vmap(potential))(r) pot = jnp.triu(pot, 1) return jnp.sum(pot) # 约束:两点坐标和等于初始值1 def constraint(x): return x[0] + x[1] - 1.0 result = minimize(F, x_init, method='SLSQP', constraints={'type': 'eq', 'fun': constraint}) x_solutions = result.x print("优化后的位置:", x_solutions) print("两点距离:", jnp.abs(x_solutions[0] - x_solutions[1])) print("理论最优距离:", 2**(1/6))
两种方法都能得到接近1.12的最优距离结果。
内容的提问来源于stack exchange,提问作者Heng Yuan
相关产品推荐
相关产品推荐

