ARTICLE DETAIL

资讯详情

深耕郑州网站建设与运营推广的一线实战洞察。

深度展开网络LISTA:让压缩感知重构告别手工调参

深度展开网络LISTA:让压缩感知重构告别手工调参 简介一套完整的深度压缩感知与学习迭代收缩阈值算法LISTA实现项目代码面向深度学习与压缩感知交叉方向的研究人员和工程师主要解决传统迭代收缩阈值算法在信号重构时计算量大、速度慢的问题。代码基于PyTorch框架编写包含两种核心算法实现原始迭代软阈值算法与可端到端训练的学习迭代收缩阈值算法并附带完整的仿真实验流程。使用者可以运行程序生成稀疏信号对比两种算法在重构精度和时间开销上的差异同时借助训练损失曲线和重构结果图片直观观察收敛过程与重构效果。压缩包共汇总9个文件主要由三个Python脚本构成核心功能代码两个PNG图片保存实验输出两个pyc文件为预编译缓存另有文本格式的依赖说明与一个项目配置文件整体大小仅为688KB结构清晰、轻量便携。已有189人学习适合希望通过代码实践深入理解深度展开优化思想、快速复现LISTA重构实验的读者。 做压缩感知重构实验的时候我把 ISTA 迭代跑了三千多次才勉强把一组 50 维的稀疏信号恢复出来。更烦的是换一个观测矩阵就得重新调步长、调阈值整套流程在工程落地里根本没法省心。后来我接触到深度展开Deep Unrolling的思路才发现迭代算法本身也可以学习——把固定的人工迭代过程换成可训练的网络层ISTA 就被改造成了 LISTALearned Iterative Shrinkage-Thresholding Algorithm。这篇文章会把 LISTA 的原理、PyTorch 实现和训练调参的完整过程拆开讲清楚适合正在接触压缩感知、稀疏重构或者深度展开网络的同学参考代码量不大但里面有几个坑确实值得提前避开。1. 从ISTA到深度展开压缩感知重构的典型痛点1.1 先回顾一下压缩感知的基本模型压缩感知处理的问题通常长这样我们有一个观测矩阵 A维度是 M×N其中 M 远小于 N。原始信号 x 是 N 维的但它在某个基底下是稀疏的或者本身就是稀疏的观测结果 y Ax。因为 M N这个方程是欠定的理论上解不唯一但借助 x 的稀疏性我们可以把它变成一个稀疏重构问题min ||x||_0 subject to y AxL0 范数直接优化是 NP 难问题所以实际中会转成 L1 范数凸优化或者用各种迭代近似算法来逼近稀疏解。经典的求解方式主要有三大类贪婪类算法OMP、CoSaMP、凸优化类LASSO、Basis Pursuit、迭代阈值类ISTA、FISTA、AMP。我自己最早用的是 OMP简单直接但抗噪性能一般而且对稀疏度 K 的估计误差很敏感。后来换到 ISTA思路倒是干净每次迭代先往梯度下降方向走一步再做一次软阈值收缩。1.2 ISTA迭代公式与那几个让人头疼的超参数ISTA 的迭代更新可以写成这样x^{k1} soft_threshold( x^k - γ A^T (A x^k - y), θ )其中 γ 是步长θ 是软阈值参数通常与正则化系数 λ 相关。soft_threshold(z, θ) sign(z) * max(|z| - θ, 0)。这个公式形式很美但工程上难受的点也集中在这两个参数上γ 的取值直接决定收敛性。如果 γ 1/||A||_2^2迭代很可能发散选小了又收敛得极慢。理论上可以用 Lipschitz 常数来定但实际中 A 的条件数不好控制几千次迭代下去误差还是慢吞吞地掉。θ 的取值则决定了解的稀疏度。θ 太大解被压得太狠信号幅度严重失真θ 太小稀疏约束形同虚设噪声也一起保留下来。所以每换一个观测矩阵、每换一组噪声水平都得重新试参。这种手调迭代超参数的模式在离线实验里还能忍一旦要端到端部署到实时系统里就很不友好了。1.3 深度展开的核心动机学习迭代而不是手调迭代深度展开这个思想最早能追溯到 Gregor 和 LeCun 在 2010 年提出的 LISTA2016 年之后随着各类 learned optimization 的工作又火了起来。核心思路其实一句话就能说明白既然迭代算法在设计上就是反复执行同一种更新规则直到收敛那这个更新规则里的线性变换和阈值参数凭什么不能是网络学习出来的换句话说ISTA 的每一次迭代其实就是一个结构固定的计算单元线性变换 → 加法 → 软阈值。如果把 T 次迭代展开成 T 层网络每一层里允许网络自己调整那些线性变换的系数和阈值参数那么网络训练收敛之后它就在隐式地建模一个针对当前数据分布的最优迭代策略。这就是 LISTA 做的事情也是它名字里 Learned 的来源。2. LISTA是如何把迭代变成网络层的2.1 从ISTA一步推导到LISTA层我们先盯着 ISTA 的迭代式看x^{k1} soft_threshold( x^k - γ A^T A x^k γ A^T y, θ )把 x^k 和 y 分别提出来这个式子本质上就是x^{k1} soft_threshold( W_e y W_g x^k, θ )其中 W_e γ A^TW_g I - γ A^T A。你看只要允许 W_e 和 W_g 自由变化而不是钉死在 A 的表达式上ISTA 就升级成了 LISTA 的一层。Gregor 和 LeCun 最开始做 LISTA 时就是基于这个观察把每层参数从 A 的固定函数变成了可学习权重矩阵。一个典型的 LISTA 层可以写为x^{k1} η_θ ( W_e y W_g x^k )其中 η_θ 是软阈值函数θ 可以是逐层的标量、逐元素的向量也可以直接作为网络参数学习。2.2 W_e、W_g和阈值θ到底在学什么理解 LISTA 的关键在于搞清楚这三组可学习参数各自的意义。W_e 承担的是从观测 y 中提取信息的角色。在原始 ISTA 里它就是 γ A^T本质上是观测矩阵的伪逆的近似。学习之后它可以自行调整把那些对稀疏重构最有判别性的观测方向给予更高权重相当于自动做了一次针对数据分布的特征筛选。W_g 承担的是残差传播的角色。在原始 ISTA 里它是 I - γ A^T A作用是从当前估计中去除掉已经被解释的成分。学习之后它可以更聪明地处理维度之间的耦合而不局限于 A^T A 给出的固定几何关系。这也解释了为什么 LISTA 往往能用很少的层数达到 ISTA 几百上千次迭代的效果——它学会了更高效的去耦合方式。θ 则是稀疏度与保真度的平衡器。传统方法里它需要人工根据噪声水平调在 LISTA 里它直接跟着反向传播一起更新。训练时如果真实信号确实稀疏网络会倾向于学到一个合理的 θ如果噪声偏大θ 也会自动变大一些。这种自适应能力正是让我觉得 LISTA 比传统算法省心的第一原因。2.3 参数共享与全参数化的取舍LISTA 在实现上有一个容易忽略的细节每一层是独立参数还是共享同一套参数共享参数版本所有层共用同一对 W_e、W_g 和 θ参数量大幅减少网络结构类似循环神经网络训练更稳定但收敛后性能略受限。全参数化版本每一层都有自己的 W_e^k、W_g^k、θ^k表达能力更强早期 LISTA 论文里采用的多是这个方案。缺点是参数随层数线性增长训练需要更多数据和更好的初始化。我个人在中等规模问题N 在 100 量级上试下来感觉层数在 10 到 15 层时全参数化比共享参数的重构精度能高 1 到 2 个 dB但训练时间大概多出 30%。如果资源紧张从共享参数版本起步、再逐步放开到独立参数是一个比较务实的路线。3. 基于PyTorch的LISTA实现全流程3.1 网络结构定义与参数初始化实现 LISTA 之前一个最重要的工程决策就是初始化。理论上直接把 W_e 初始化为 γ A^TW_g 初始化为 I - γ A^T A网络在训练开始时就会退化成 T 层 ISTA。这种做法能保证训练起点不差后续学习只是在 ISTA 的基础上做微调稳定性比随机初始化好太多。我试过用随机高斯初始化训练相同结构结果好几次都收敛到很差的局部最优。下面是完整的网络定义我用了一个共享参数的版本便于理解import torch import torch.nn as nn import numpy as np class LISTA(nn.Module): def __init__(self, A, T12, threshold_init0.01): super(LISTA, self).__init__() # A: (M, N) 观测矩阵M表示观测数N表示信号长度 M, N A.shape self.T T self.N N # 计算ISTA步长的Lipschitz上界用于初始化 gamma 1.0 / torch.linalg.norm(A A.T, 2) I torch.eye(N) # 初始化策略让网络起点等价于T层ISTA We_init gamma * A.T # (N, M) Wg_init I - gamma * A.T A # (N, N) self.We nn.Parameter(We_init.clone()) self.Wg nn.Parameter(Wg_init.clone()) self.theta nn.Parameter(torch.full((N,), threshold_init)) # 这里用逐元素阈值向量比标量阈值更灵活 def forward(self, y): # y: (batch, M) x torch.zeros(y.shape[0], self.N, devicey.device) for _ in range(self.T): z y self.We.T x self.Wg.T x self.soft_threshold(z) return x def soft_threshold(self, z): return torch.sign(z) * torch.clamp(torch.abs(z) - self.theta, min0.0)注意这里我用了逐元素的阈值向量而不是标量。这样网络可以在训练中针对不同信号维度学出不同的收缩强度效果会比单一标量阈值更细腻。如果你希望进一步限制模型复杂度把nn.Parameter(torch.full((N,), threshold_init))换成nn.Parameter(torch.tensor(threshold_init))即可。3.2 数据生成稀疏信号、固定观测矩阵与噪声训练 LISTA 不像训练普通分类网络那样需要大规模数据集关键在于生成的数据要尽量贴近实际使用场景。我通常的做法是固定一个观测矩阵 A随机生成大量稀疏信号然后计算带噪观测 y。def generate_data(num_samples, M, N, K, noise_level0.01, seed0): rng np.random.default_rng(seed) A torch.tensor(rng.normal(0, 1, (M, N)) / np.sqrt(M), dtypetorch.float32) signals torch.zeros(num_samples, N) for i in range(num_samples): idx rng.choice(N, K, replaceFalse) signals[i, idx] rng.normal(0, 1, K) y signals A.T if noise_level 0: y torch.randn_like(y) * noise_level return A, signals, y这里的稀疏度 K 是关键先验我建议训练用的 K 与实际场景保持一致。如果训练时 K5测试时遇到 K15 的信号重构误差会明显上升。另外观测矩阵 A 固定不变时网络学习的是针对这个 A 的最优迭代策略如果 A 也随样本变化那模型参数量就很难覆盖所有情况这点在后面章节还会细聊。3.3 训练闭环与损失函数选择训练部分比较常规但有一个细节值得强调我倾向于用带噪声的观测数据训练因为纯净的无噪情形下网络很容易过度拟合观测矩阵的代数结构测试时一遇到噪声就崩。def train_lista(model, A, train_signals, train_y, epochs150, lr1e-3, batch_size256): optimizer torch.optim.Adam(model.parameters(), lrlr) scheduler torch.optim.lr_scheduler.CosineAnnealingLR(optimizer, T_maxepochs) loss_fn nn.MSELoss() dataset torch.utils.data.TensorDataset(train_y, train_signals) loader torch.utils.data.DataLoader(dataset, batch_sizebatch_size, shuffleTrue) for epoch in range(epochs): model.train() total_loss 0.0 for y_batch, x_batch in loader: optimizer.zero_grad() x_hat model(y_batch) loss loss_fn(x_hat, x_batch) loss.backward() optimizer.step() total_loss loss.item() * y_batch.shape[0] scheduler.step() if (epoch 1) % 30 0: print(fEpoch {epoch1:3d}, Loss: {total_loss / len(dataset):.6f})损失函数用单纯的 MSE 就够用。有人会问要不要在损失里加 L1 稀疏正则我的实测结论是如果标签本来就是稀疏信号MSE 已经隐式地推动网络输出接近稀疏信号额外加 L1 反而会让阈值偏保守重构结果幅度被压低。训练结束后测试评估可以用两个指标归一化均方误差NMSE和支撑集恢复准确率。NMSE 定义是 ||x - x_hat||_2^2 / ||x||_2^2支撑集准确率则是估计稀疏位置与真实稀疏位置重叠的比例。两者结合看能同时反映幅度准确性和稀疏模式恢复能力。4. 训练LISTA绕不开的几个坑与调参经验4.1 初始化不是锦上添花而是成败关键我前面提到过LISTA 的初始化决定了训练起点。这里我再展开说一个具体踩坑案例。有一组实验里我把 W_e 和 W_g 都换成了标准正态随机初始化训练到第 80 轮时损失不降反升最后重构结果甚至不如固定参数的 ISTA。后来我复盘原因在于随机初始化让网络一开始就失去了 ISTA 的先验结构——它的迭代行为完全是乱的梯度信号在这种随机往返过程中被稀释了。所以我的建议很明确如果 W_e、W_g 没有更好的先验就直接用 γA^T 和 I - γA^TA 初始化。这种初始化相当于在专家设计的迭代基础上局部搜索训练效率和安全系数都高很多。另外这里的 γ 最好用1 / ||A||_2^2的精确值而不是粗略估计我用torch.linalg.norm(A A.T, 2)计算谱范数的平方倒数实测很稳。4.2 网络深度不是越深越好阈值向量要防崩塌LISTA 的层数 T 是一个需要平衡的超参数。层数越多网络对迭代的展开越充分理论上限越高但参数量和过拟合风险也同步上升。在我常用的 50 维信号、30 维观测场景下T 从 8 涨到 12 时重构误差有明显下降但 T 继续涨到 20 以上后收益就很小了反而训练后期偶尔出现震荡。如果你的数据量不大建议从 T10 起步逐步加层观察验证集误差。阈值向量也有个常见坑如果初始化阈值太大而训练数据噪声较低网络可能直接把阈值向量推到接近 0这时候软阈值函数退化成恒等函数稀疏约束完全失效。为了避免这种阈值崩塌我习惯给阈值参数加一个很小的权重衰减或者用torch.clamp在 forward 里把 θ 限制在 [1e-4, 0.5] 范围内。此外训练初期如果 loss 下降极慢优先检查一下 θ 的绝对值是不是已经变得不合理的巨大。4.3 与 ISTA 和 OMP 的对比结果用我自己的实验数据来给个直观参照。实验设置是N50M30K5噪声水平 0.01观测矩阵 A 固定。三种方法对比结果如下固定参数的 ISTA迭代 500 次NMSE 约 0.13支撑集准确率约 92%。OMP迭代 5 次K 已知NMSE 约 0.09支撑集准确率约 94%。LISTAT12训练完成后推理只做 12 层矩阵乘加软阈值NMSE 约 0.04支撑集准确率约 98%。LISTA 的推理速度比 500 次迭代的 ISTA 快了一个数量级以上因为 500 次迭代要反复做矩阵乘法而 12 层 LISTA 本质上就是 12 次矩阵乘加一个元素级非线性。这同时也是深度展开网络最大的工程价值它把时间换精度的迭代过程压缩成了一次性前向传播。4.4 观测矩阵变化的场景与迁移应对最后说一个 LISTA 在工程应用里最容易被忽视的边界如果观测矩阵 A 在运行时发生了变化已经训练好的 LISTA 权重不会自动适配。比如在传感矩阵做了校准调整后同样的输入 y网络输出误差会突然变大。这种场景下我没有在模型结构里引入特殊机制而是采取了一个很实用的方案直接将 W_e、W_g 中与 A 相关的部分重新初始化然后做几十轮精细微调fine-tune。因为网络结构没变微调不需要重新生成大量数据往往几百个样本就能收敛回可接受水平。这种做法比重新训练整个模型效率高得多也是我个人在项目中最常用的兜底手段。如果你确定未来观测矩阵会持续变化更长期的办法是使用结构化 LISTA比如让 W_e 和 W_g 在 forward 时共享 A 的某种参数化表示而不是直接学习独立矩阵。但这样实现复杂度会上一个台阶适用面也没有标准 LISTA 广。我的建议是先把手头固定 A 的场景做好真的遇到变化时用 fine-tune 过渡不要一开始就过度设计。在实际训练中我还发现一个小技巧训练数据里的稀疏信号支撑集位置如果总是固定在某些维度模型会对这些维度形成偏好。所以每次生成数据时我都用不同的随机种子重新采样支撑集位置确保网络学会的是通用的稀疏重构能力而不是记住某几个坐标。这个习惯后来帮我避免了好几次实验误差。本文还有配套的精品资源点击获取
返回列表