如何用Jax/JaxOpt求解线性规划?现有尝试失败求排查
JaxOpt求解最优运输LP问题的错误排查与解决方案
常见错误原因分析
- 数值稳定性缺失:最优运输的成本矩阵常存在极端数值,Jax自动微分对这类值敏感,易引发NaN;将LP转为QP时,若二次项系数设置过接近0,会导致Hessian矩阵病态,求解器直接发散。
- 约束转换错误:最优运输的核心是等式边际约束(行和=源分布、列和=目标分布),若转QP时误设为不等式,或约束矩阵维度对齐错误,会直接导致求解结果偏离可行域。
- 求解器参数适配差:JaxOpt QP求解器的默认迭代次数、步长、正则化参数不匹配最优运输问题结构,迭代未收敛就终止,出现异常极值或NaN。
- 初始值偏离可行域:若初始值完全不符合边际约束(比如全0随机值),求解器迭代过程中会进入数值不稳定区域,无法收敛到合理解。
针对性解决方案
1. 正确实现LP转QP的形式
最优运输LP转QP时,需给Hessian矩阵添加极小正则化项保证正定性,同时严格对齐约束维度:
import jax import jax.numpy as jnp from jaxopt import QP # 假设C为成本矩阵,a为源分布,b为目标分布 n, m = C.shape # 展开运输矩阵P为向量 flatten_P = lambda P: P.reshape(-1) # QP目标函数:添加极小正则化项避免Hessian奇异 H = jnp.eye(n*m) * 1e-7 c = flatten_P(C) # 构建等式约束矩阵:源边际+目标边际 A_eq_source = jnp.kron(jnp.eye(n), jnp.ones((1, m))) # 行和约束 A_eq_target = jnp.kron(jnp.ones((1, n)), jnp.eye(m)) # 列和约束 A_eq = jnp.vstack([A_eq_source, A_eq_target]) b_eq = jnp.hstack([a, b]) # 非负约束:-P <= 0 A_ineq = -jnp.eye(n*m) b_ineq = jnp.zeros(n*m) # 初始化求解器,调整参数保证收敛 qp_solver = QP(verbose=True, maxiter=1500, tol=1e-6) # 初始值选用符合边际约束的独立分布(避免偏离可行域) init_P_flat = flatten_P(jnp.outer(a, b)) sol = qp_solver.run(init_P_flat, H=H, c=c, A_eq=A_eq, b_eq=b_eq, A_ineq=A_ineq, b_ineq=b_ineq) # 恢复为运输矩阵 P_opt = sol.params.reshape(n, m)
2. 提升数值稳定性
- 对成本矩阵
C做归一化处理:比如C = C / jnp.max(C),避免极端数值引发的微分溢出。 - 调试正则化项
H的取值:建议在1e-8到1e-6之间选择,平衡原LP问题的精度和求解器稳定性。
3. 适配求解器参数与投影机制
若使用梯度下降类求解器,必须添加投影操作保证每一步迭代都在可行域内:
from jaxopt import GradientDescent from jaxopt.projection import projection_non_negative, projection_affine # 定义最优运输目标函数 def objective(P): return jnp.sum(C * P) # 组合投影:先非负投影,再边际约束投影 def projection(P): P_non_neg = projection_non_negative(P) return projection_affine(P_non_neg, A_eq, b_eq) # 初始化梯度下降求解器,调整步长与迭代次数 gd_solver = GradientDescent(fun=objective, projection=projection, maxiter=2000, step_size=0.05) init_P = jnp.outer(a, b) sol = gd_solver.run(init_P)
4. 验证约束正确性
- 检查
A_eq维度:行数应为n+m(n个源约束+m个目标约束),列数为n*m(展开后的P向量长度)。 - 验证初始值:
jnp.sum(init_P, axis=1)应接近a,jnp.sum(init_P, axis=0)应接近b,确保初始点在可行域附近。
内容的提问来源于stack exchange,提问作者logan
相关产品推荐
相关产品推荐

