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

基于PyTorch GPU求解带非负约束的大型Ax=b线性系统

最优解决方案

你的问题本质是非负最小二乘(NNLS)问题(当Ax=b无解时,找x≥0使||Ax - b||²最小;有解时则是该问题的特例)。结合你无法存储矩阵A、仅能高效计算Ax的需求,以下是基于PyTorch GPU加速的最优实现方案:

1. 投影梯度下降(基础高效版)

梯度下降仅依赖目标函数的梯度计算,而目标函数L(x)=0.5*||Ax - b||²的梯度为∇L(x)=Aᵀ(Ax - b)。对于无法直接获取A的场景,可通过PyTorch自动微分间接计算Aᵀv(v为任意向量),再配合非负投影完成迭代:

import torch

def compute_Ax(x):
    # 替换为你的高效Ax计算逻辑(如卷积、自定义线性变换等)
    # 示例:假设A是大尺寸的线性变换,通过自定义逻辑计算Ax
    return custom_efficient_transform(x)

def compute_A_T_v(v):
    # 利用自动微分计算A^T v
    x = torch.zeros_like(v, requires_grad=True, device=v.device)
    Ax = compute_Ax(x)
    torch.sum(Ax * v).backward()
    return x.grad

def nnls_gradient_descent(b, x_init, lr=1e-3, num_iter=1000, tol=1e-6):
    x = x_init.clone().detach().to(b.device)
    for _ in range(num_iter):
        Ax = compute_Ax(x)
        residual = Ax - b
        grad = compute_A_T_v(residual)
        # 梯度更新后投影到非负空间
        x = x - lr * grad
        x = torch.clamp(x, min=0.0)
        # 收敛判断
        if torch.norm(residual) < tol:
            break
    return x
  • 优势:实现简单,完全兼容黑箱Ax计算,PyTorch自动将计算逻辑映射到GPU,加速效果显著。
  • 优化点:可替换SGD为Adam、Adagrad等自适应优化器,提升收敛稳定性。

2. FISTA快速迭代算法(收敛加速版)

针对梯度下降收敛慢的问题,采用FISTA(快速迭代收缩阈值算法),它在凸问题上具有更快的收敛速率,核心是引入动量项优化迭代过程:

def nnls_fista(b, x_init, lr=1e-3, num_iter=1000, tol=1e-6):
    x = x_init.clone().detach().to(b.device)
    y = x.clone()
    t = 1.0
    for _ in range(num_iter):
        Ax_y = compute_Ax(y)
        residual = Ax_y - b
        grad = compute_A_T_v(residual)
        x_new = torch.clamp(y - lr * grad, min=0.0)
        # 更新动量参数
        t_new = (1 + torch.sqrt(1 + 4 * t**2)) / 2
        y = x_new + ((t - 1)/t_new) * (x_new - x)
        x = x_new
        # 收敛判断
        if torch.norm(compute_Ax(x) - b) < tol:
            break
    return x

3. PyTorch优化器封装版(最简实现)

直接利用PyTorch内置优化器,通过自动微分处理梯度计算,仅需在每步优化后手动施加非负约束:

def nnls_pytorch_optimizer(b, x_init, lr=1e-3, num_iter=1000, tol=1e-6):
    x = x_init.clone().detach().requires_grad_(True).to(b.device)
    optimizer = torch.optim.Adam([x], lr=lr)
    for _ in range(num_iter):
        optimizer.zero_grad()
        Ax = compute_Ax(x)
        loss = torch.norm(Ax - b)**2 / 2
        loss.backward()
        optimizer.step()
        # 投影到非负空间(禁用梯度追踪)
        with torch.no_grad():
            x.clamp_(min=0.0)
        if loss.item() < tol:
            break
    return x.detach()
  • 优势:无需手动实现梯度和迭代逻辑,借助PyTorch成熟的优化器生态,开发效率极高。

关键注意事项

  • GPU适配:确保所有张量(x、b)通过.to('cuda')移至GPU,compute_Ax的逻辑需兼容PyTorch GPU张量(原生算子自动支持,自定义CUDA核需实现反向传播)。
  • 收敛性:上述方法均针对凸问题,可保证收敛到全局最优解,优先推荐FISTA以平衡收敛速度和实现复杂度。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.18 03:01:22