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

Coordinate Descent算法实现不收敛问题求助(附最小可复现代码)

Coordinate Descent算法实现不收敛问题求助(附最小可复现代码)

我正在尝试实现论文《Learning Fast Approximations of Sparse Coding》中描述的算法,已经完成了ISTA、FISTA算法以及基于投影块坐标下降的字典更新规则。

不过在实现用于近似最优稀疏编码的**Coordinate Descent(论文中的Algorithm 2)**时遇到了问题——代码似乎无法收敛。之前ISTA也出现过类似问题,但很快我就发现是学习率小于$D^T D$的最大特征值,违反了收敛条件。但Coordinate Descent算法除了正则化参数外没有这类超参数,这让我很困惑。

我试过调整正则化参数,但完全没有效果。

最小可复现代码

import torch
import itertools

def shrink(a, b):
    return torch.sign(a) * torch.maximum(torch.abs(a) - b, torch.tensor(0.0))

def change(a, b):
    return torch.sum(torch.abs(a - b))

def CoD(x : torch.Tensor,
         h_dim : int,
         D : torch.Tensor,
         regularization=0.5,
         frequency=None) -> torch.Tensor:
    """
    Implements the Coordinate Descent algorithm to find the optimal
    sparse codes for the given input vector x and a given
    dictionary matrix D. This function follows notation corresponding to
    the paper "Learning Fast Approximations of Sparse Coding"

    :param x: The input vector
    :param h_dim: Dimension of sparse code required. In
    the overcomplete case, should be greater than or equal to
    dim(x).
    :param D: the dictionary matrix
    :param regularization: parameter which controls the relative
    weight given to sparsity vs reconstruction loss
    :param lr: the learning rate
    :param frequency: the number of iterations after which an update is printed
    :return: An approximation to the optimal sparse code of the
    given input h_opt.
    """
    n = x.shape[0]
    m = h_dim
    assert D.shape == (n, m), (f"D should be matrix of shape (x_dim, h_dim), where "
                               f"x_dim={n} and h_dim={m}, but dimensions of D are {D.shape}")
    z = torch.zeros(h_dim)
    z_old = z.clone()

    S = torch.eye(h_dim) - torch.matmul(torch.transpose(D, 0, 1), D)
    B = torch.matmul(D.T, x)

    counter = itertools.count(start=0, step=1)
    while True:
        z = shrink(B, regularization)
        k = torch.argmax(torch.abs(z - z_old))
        for j in range(m):
            B[j] = B[j] + S[j][k] * (z[k] - z_old[k])
        if change(z, z_old) < 0.01:
            break
        z_old = z.clone()
        iter_num = next(counter)
        if frequency is not None and iter_num % frequency == 0:
            print(f"Coordinate Descent: Iteration {iter_num}")
    return shrink(B, regularization)

print(CoD(torch.randn(10), 10, torch.randn(10, 10), frequency=100))

我几乎完全按照论文描述来实现算法,有人能帮我调试一下吗?

我检查了日志,发现$Z, \bar{Z}$和B的值会呈指数级增长,直到变成无穷大,最后变成NaN。

备注:内容来源于stack exchange,提问作者insipidintegrator

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.04.14 16:52:59