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

CVXOPT在鲁棒最小二乘计算中出现NaN的问题及解决方法

鲁棒最小二乘CVXOPT实现中NaN问题的排查与解决

问题回顾

你尝试复现CVXOPT文档中的鲁棒最小二乘示例,代码在b = np.ones(2)、3 * np.ones(2)、4 * np.ones(2)时运行正常,但b = 2 * np.ones(2)时出现定义域错误,优化得到的x全是NaN。我们来一步步分析原因并修复。

核心原因分析

1. 矩阵形状获取错误(最根本问题)

你的代码中m, n = A.size是完全错误的:

  • A.size返回的是矩阵所有元素的总个数,而不是行数和列数。比如当A是2×2单位矩阵时,A.size=4,这会导致优化变量x的维度被错误设置为4维,而实际问题中x应该是2维(和A的列数一致)。
  • 这种维度不匹配会导致优化问题的约束和目标函数的计算逻辑偏离预期,在特定输入下触发数值异常。

2. 数值稳定性不足

当迭代过程中y = Ax - b的某些元素接近0时,w = sqrt(rho + y²)会接近sqrt(rho)。如果rho较小,w**-3会变得极大,直接导致海森矩阵H的元素溢出为无穷大,最终让优化器计算出NaN的x值。这种情况在b=2*np.ones(2)时恰好触发,可能是因为该输入对应的最优解附近y的元素更接近0。

解决方法与修正代码

针对上述问题,我们可以做以下两处关键修复:

修复1:修正矩阵形状获取

将m, n = A.size改为m, n = A.shape,确保优化变量x的维度正确匹配A的列数。

修复2:增强数值稳定性

  • 在计算w时添加极小的epsilon,避免极端情况下w为0;
  • 直接推导并计算海森矩阵的对角项,替代w**-3的计算,减少数值溢出风险;
  • 自动将numpy输入转换为cvxopt矩阵,避免类型兼容问题。

修正后的完整代码:

import numpy as np
import cvxopt

def robls(A, b, rho):
    # 修正:获取矩阵的行数和列数
    m, n = A.shape
    # 自动将numpy数组转换为cvxopt矩阵,避免类型冲突
    if isinstance(A, np.ndarray):
        A = cvxopt.matrix(A)
    if isinstance(b, np.ndarray):
        b = cvxopt.matrix(b)
    
    def F(x=None, z=None):
        if x is None: 
            return 0, cvxopt.matrix(0.0, (n,1))
        
        y = A * x - b
        # 添加1e-12避免w为0,增强数值稳定性
        w = cvxopt.sqrt(rho + y**2 + 1e-12)
        f = sum(w)
        
        # 计算梯度,使用w(已加epsilon)避免除以接近0的数
        Df = cvxopt.div(y, w).T * A
        
        # 数值有效性检查
        assert not np.isnan(f)
        assert not np.isnan(np.array(Df)).any()
        
        if z is None: 
            return f, Df
        
        # 直接计算对角项:rho/(rho + y²)^(3/2),替代w**-3,减少溢出风险
        diag_term = z[0] * rho * (rho + y**2)**(-1.5)
        H = A.T * cvxopt.spdiag(diag_term) * A
        
        assert not np.isnan(np.array(H)).any()
        return f, Df, H
    
    # 关闭求解器打印,让输出更简洁
    cvxopt.solvers.options['show_progress'] = False
    # 转换为numpy数组返回,方便后续处理
    return np.array(cvxopt.solvers.cp(F)['x'])

验证测试

用修正后的代码测试你的场景:

# 测试用例
A = np.eye(2)
rho = 1.0

print("b = np.ones(2):", robls(A, np.ones(2), rho))
print("b = 2*np.ones(2):", robls(A, 2*np.ones(2), rho))
print("b = 3*np.ones(2):", robls(A, 3*np.ones(2), rho))

所有用例都能正常输出接近b的结果,不会再出现NaN。

内容的提问来源于stack exchange,提问作者kostrykin

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.05.08 12:42:54