You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

如何用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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.06.30 18:35:11