Pruning Large Language Models with Semi-Structural Adaptive Sparse Traini 论文精读:让大语言模型在训练中主动适应 2:4 稀疏

发布时间:2026/7/24 20:55:26

Pruning Large Language Models with Semi-Structural Adaptive Sparse Traini 论文精读:让大语言模型在训练中主动适应 2:4 稀疏 论文提出的方法叫作Adaptive Sparse Trainer简称 AST。其核心问题是与其先确定一个固定的 2:4 剪枝掩码再训练剩余权重能否在训练过程中同时调整权重和稀疏掩码让模型自己找到更合适的半结构化稀疏连接一、论文基本信息项目内容论文题目Pruning Large Language Models with Semi-Structural Adaptive Sparse Training方法名称Adaptive Sparse TrainerAST作者Weiyu Huang、Yuezhou Hu、Guohao Jian、Jun Zhu、Jianfei Chen会议AAAI 2025发表时间2025年4月页码24167–24175DOI10.1609/aaai.v39i23.34592论文链接AAAI正式版本、arXiv扩展版本官方代码thu-ml/Adaptive-Sparse-Trainer论文正式发表于AAAI 2025官方代码提供了GPT-2、OPT和LLaMA-2-7B相关实现。二、论文要解决什么问题2.1 一次性2:4剪枝容易破坏模型能力SparseGPT和Wanda可以使用少量校准数据快速将大语言模型剪成2:4稀疏模型。但是2:4不是任意删除一半权重而是要求在每一组连续4个权重中必须恰好保留2个、删除2个。例如某一组权重为[0.8-0.10.50.2]2:4剪枝后只能保留其中两个。按照幅值可能得到[0.800.50]这种局部约束有利于硬件执行但限制远强于普通非结构化剪枝。即使模型整体存在大量冗余也不代表每个连续4权重的小组都能安全删除一半。论文指出一次性剪枝方法在困惑度上可能还可以接受但在MMLU、数学推理和其他知识密集任务上性能损失通常更加明显。2.2 固定剪枝掩码会限制恢复能力一种常见做法是先用Wanda或SparseGPT确定掩码 → 固定掩码 → 只训练剩余权重这种方法的问题在于初始剪枝判断一旦出错就很难再修正。假设某组权重中原本删除了A和B但经过训练后发现A实际上对模型很重要当前保留的C作用较小。如果掩码已经固定A无法重新进入网络C也不能被替换。后续训练只能在一个可能并不理想的稀疏子网络中继续优化。论文认为对于严格的2:4稀疏只学习权重而不学习连接位置往往无法获得最优稀疏模型。2.3 普通稀疏训练存在掩码振荡问题另一种做法是在训练过程中不断根据权重幅值重新选择每组保留的两个权重。这样虽然允许掩码变化却容易出现另一个问题某个权重这一步进入掩码下一步又被删除再下一步重新进入不同权重频繁交换位置。如果掩码一直剧烈振荡模型无法稳定收敛。因此动态稀疏训练需要同时解决两个相反目标前期要允许掩码充分探索后期又必须让掩码逐渐稳定。2.4 稀疏模型重新训练容易陷入局部最优作者还观察到一个现象称为Retraining Dilemma即重新训练困境。稀疏模型继承了预训练模型的大量已有知识因此在重新训练开始时训练损失往往下降得很快。但与此同时测试集困惑度可能长期不稳定甚至明显高于从头训练的模型。作者认为模型会快速拟合有限的重新训练数据却没有充分学习适合当前稀疏结构的全局表示。单纯降低学习率也不能完全解决这一问题。三、半结构化稀疏到底是什么大语言模型剪枝通常可以分为三种粒度。类型删除对象优点主要问题非结构化剪枝任意单个权重精度保持较好稀疏位置不规则硬件难加速结构化剪枝神经元、通道、注意力头、层容易获得真实加速删除粒度过大精度容易下降半结构化剪枝每个局部权重组保留固定数量兼顾灵活性与硬件规则性局部约束严格掩码较难优化AST主要研究2:4半结构化稀疏。它要求在线性层每一行中每连续4个权重只能有2个非零值整体权重稀疏率固定为50%。与普通50%非结构化剪枝相比两者虽然零权重比例相同但后者可以自由选择模型中最不重要的一半参数2:4则必须在每个局部小组中删除一半因此通常更难保持精度。2:4的优势在于规则模式可以被Tensor Core和TensorRT-LLM等软件栈利用在支持的NVIDIA GPU上减少矩阵乘法及权重读取开销。四、AST的核心思想AST由三个主要模块组成模块作用Annealing SR-STE训练过程中动态调整2:4掩码并逐渐稳定下来知识蒸馏使用稠密模型指导稀疏模型缓解重新训练困境SLoRB使用被删除权重初始化一个额外的低秩补偿分支其中前两个构成基本的AST-NaiveSLoRB是可选增强模块加入后称为AST-Boosted。整个思路可以概括为训练初期允许被剪权重重新进入模型让网络探索不同的2:4连接方式训练后期逐渐增强衰减使不重要权重稳定接近零同时利用稠密教师的输出分布帮助稀疏模型保持原有知识。五、为什么不能直接固定掩码固定掩码方法通常执行以下流程使用幅值、Wanda或SparseGPT确定2:4掩码将被删除权重设置为零后续训练只更新保留权重被删除权重永远不能恢复。AST则保留一份完整的底层权重。在前向传播时模型仍然按照2:4掩码执行稀疏计算但训练过程中当前未被选中的权重仍然可以接收梯度。每隔一定训练步数AST重新按照当前权重幅值生成2:4掩码。因此一个当前被删除的权重如果持续获得较强梯度其绝对值可能重新增大并在下一次掩码更新时替换掉其他权重。这个过程可以理解为删除不是永久判决而是暂时退出竞争。直到训练后期掩码逐渐稳定最终才形成用于部署的固定2:4结构。论文的消融实验表明动态掩码比从训练开始就固定掩码表现更好。六、Annealing SR-STE如何同时实现探索和稳定6.1 STE让离散掩码可以参与训练2:4掩码本质上是离散选择保留权重对应1删除权重对应0。这种排序和二值化操作不能直接计算普通梯度。AST使用直通估计器STE处理这个问题。直观上看前向传播使用真正的2:4稀疏权重反向传播时近似忽略掩码不可导的问题梯度可以继续传回底层权重。这样当前被掩蔽的权重仍有机会在训练中改变数值。6.2 为什么要对被掩蔽权重施加衰减如果所有权重都只按照任务梯度更新当前被删除和保留的权重可能长期处在相近幅值不同位置会频繁交换导致掩码震荡。AST因此对当前未被选中的权重施加额外的L2衰减使不重要权重逐渐趋近于零。但衰减强度非常关键衰减过强被删除权重几乎无法重新进入动态掩码退化成固定掩码衰减过弱大量权重不断交换位置模型难以收敛。6.3 退火衰减策略AST提出Annealing SR-STE也就是带退火过程的SR-STE。训练初期衰减系数很小被删除权重不会马上被压到零模型可以尝试不同连接组合。随着训练进行衰减逐渐增大不重要权重越来越难重新进入掩码开始稳定。达到预设阶段后衰减强度不再增加模型在较稳定的2:4结构上完成最后收敛。因此这个过程类似于前期探索掩码 → 中期筛选连接 → 后期固定结构论文对掩码翻转率进行了分析。相比使用固定衰减强度的SR-STE退火版本在训练早期改变了更多连接但在训练末期表现得更加稳定。七、AST如何选择每组保留的权重AST最终使用的是非常简单的权重幅值准则。也就是说在每组连续4个权重中保留当前绝对值最大的两个。作者也尝试过包含激活信息的Wanda或其他一次性剪枝指标但发现它们在动态训练场景中并没有优势。原因在于训练过程中权重不断更新内部激活分布也不断变化激活统计容易过时需要周期性重新收集校准激活计算成本更高后期掩码选择可能受到激活波动影响。相比之下权重幅值可以直接从当前参数获得计算开销低也更加稳定。论文报告在AST的动态训练环境中幅值准则在效率和最终性能上都优于所测试的激活型指标。这也说明适合一次性剪枝的指标不一定适合动态稀疏训练。八、知识蒸馏如何缓解重新训练困境AST将原始稠密模型作为教师将2:4稀疏模型作为学生。训练目标由两部分组成学生对真实训练文本的语言建模损失学生输出概率与稠密教师输出概率之间的KL散度。可以用纯文本表示为总损失 任务交叉熵 教师与学生输出分布差异教师提供的不只是正确Token还包含其他Token的相对概率。例如正确答案是“Paris”教师可能给出Paris0.80London0.08Berlin0.05其他0.07。普通交叉熵只告诉学生“Paris是正确答案”教师分布则同时提供了候选Token之间的相对关系。对于训练数据有限的稀疏恢复过程这种软监督比单一标签更加丰富。论文发现KL蒸馏能够减少稀疏模型早期过拟合使测试困惑度更稳定并在固定训练预算下加快收敛。8.1 为什么没有使用中间特征蒸馏上一篇Sparse Fine-tuning论文的核心是SquareHead即要求稀疏学生匹配教师每一层的中间特征。但AST得出了不同结论。作者测试了TinyBERT、MobileBERT和Sparse Fine-tuning等包含隐藏状态或注意力蒸馏的方法发现这些中间约束会降低生成式语言模型的泛化效果。仅使用输出概率的KL蒸馏已经能够获得最好的整体表现。在GPT-2-124M上相关困惑度结果为蒸馏方式困惑度TinyBERT42.75MobileBERT44.87Sparse Fine-tuning式蒸馏41.19MiniLLM反向KL32.20AST前向KL32.23前向KL和反向KL结果非常接近而中间特征蒸馏明显更差。这里并不能简单得出“中间特征蒸馏始终无效”的结论。更合理的理解是Sparse Fine-tuning主要研究任务微调后的模型恢复AST主要使用通用预训练语料继续训练生成式模型两种训练数据、模型和目标并不完全一致中间表示的强约束可能限制AST学生重新适应新的稀疏连接。九、SLoRB用被删除权重补偿模型容量9.1 为什么需要额外补偿严格2:4稀疏意味着每个线性层只有一半原始权重参与计算。即使掩码选择得很好模型表达能力仍然会下降。因此作者设计了一个可选模块Sparse Low-Rank BoostingSLoRB。它类似于LoRA在原始稀疏线性层之外增加一个低秩分支。但SLoRB和普通LoRA存在两个关键区别普通LoRASLoRB通常冻结基础模型稀疏基础权重和低秩分支同时训练低秩参数通常随机初始化使用被删除权重的信息初始化主要用于下游任务适配用于补偿2:4剪枝损失9.2 如何利用被删除权重初始化SLoRB把一段连续权重划分为更大的组然后计算该组中被删除权重的平均值用它初始化对应的低秩补偿参数。直观来看2:4剪枝删除了一半连接但这些被删除权重并不是完全没有信息。SLoRB不再逐个保存它们而是用一个较小的低秩结构保存其粗略的平均影响。这相当于主要信息由2:4稀疏权重保留被删除权重的低频或平均信息由低秩支路补充。作者认为这种初始化比随机初始化收敛更快尤其适合训练Token预算有限的场景。9.3 SLoRB并不是免费提升论文在主要实验中令SLoRB参数k等于16这会增加约12.5%的参数。因此AST-Naive是严格的2:4模型约有3.4B非零参数AST-Boosted加入SLoRB后约有4.2B参数后者精度更高但压缩率和部署收益会有所下降。所以AST-Boosted不应被理解为“完全相同模型大小下无代价提升”。它本质上使用了额外参数换取性能。十、完整训练流程AST的执行过程可以概括为第一步加载稠密预训练模型稠密模型同时充当稀疏学生的初始化知识蒸馏教师。第二步初始化2:4掩码在每组连续4个权重中暂时保留幅值最大的两个。第三步执行稀疏前向传播学生模型始终按照当前2:4掩码计算。第四步计算训练损失同时计算语言建模交叉熵和教师—学生KL蒸馏损失。第五步更新全部底层权重通过STE让当前被掩蔽权重也能够获得更新机会。第六步衰减当前未被选择的权重衰减强度在训练前期较小随后逐渐增强。第七步周期性重新生成掩码根据最新权重幅值在每个4权重小组中重新选择两个保留位置。第八步可选训练SLoRB低秩补偿分支和稀疏主干共同更新。第九步训练结束后固定掩码输出可供TensorRT-LLM等运行时使用的2:4稀疏模型。官方算法中掩码按固定间隔更新梯度来自蒸馏损失SLoRB则作为可选模块与稀疏权重共同优化。十一、实验设置论文测试了三个模型家族模型家族规模OPT125M、350M、1.3BGPT-2124M、350M、774M、1.5BLLaMA-27B较小的OPT和GPT-2模型使用C4训练LLaMA-2-7B使用RedPajama-v1其数据覆盖CommonCrawl、C4、GitHub、Wikipedia、Books、ArXiv和StackExchange等领域。(AAAI Publications)主要训练和评估设置如下设置内容稀疏形式2:4半结构化稀疏语言建模评估WikiText-2困惑度零样本任务BoolQ、RTE、HellaSwag、WinoGrande、ARC-e、ARC-c、OpenBookQA知识任务MMLU、MATH、GSM8K稠密教师原始预训练模型蒸馏方式输出Logits的KL散度速度测试TensorRT-LLMGPURTX 4090、NVIDIA L20量化测试AWQ附录表格显示GPT-2和OPT根据规模使用约2.5B至10B训练TokenLLaMA-2-7B使用约7.5B Token正文则写为7B因此论文内部存在轻微表述差异。两者都大约相当于LLaMA-2预训练Token数量的0.4%以内。十二、语言建模结果下面是2:4稀疏模型在WikiText-2上的部分结果困惑度越低越好。模型稠密SparseGPTWandaAST-NaiveAST-BoostedOPT-125M27.7645.5860.9130.2228.68OPT-350M22.0040.3350.1624.6524.03OPT-1.3B14.6229.0323.9215.8515.43GPT-2-124M29.9550.09115.6432.2331.13GPT-2-350M21.7231.0363.7123.6523.03GPT-2-774M19.4325.9849.9721.2920.66GPT-2-1.5B17.4021.1430.4418.3318.01结果说明一次性2:4剪枝在小模型上尤其容易崩溃。例如GPT-2-124M稠密模型29.95SparseGPT50.09Wanda115.64AST-Naive32.23AST-Boosted31.13。AST基本恢复了大部分困惑度损失。随着模型增大2:4剪枝相对更容易恢复。例如GPT-2-1.5B的AST-Boosted仅从17.40上升到18.01。说明较大模型拥有更多冗余和恢复空间。十三、LLaMA-2-7B零样本结果七个零样本任务的平均准确率如下模型非零参数量平均准确率LLaMA-2-7B稠密模型6.7B59.78Wanda LLaMA-7B 2:43.4B56.19AST-Naive3.4B57.68AST-Boosted4.2B58.62AST-Naive与稠密模型仍相差2.10个百分点而加入SLoRB后差距缩小到1.16个百分点。论文摘要所说的“零样本准确率差距只有1.16%”对应的是AST-Boosted不是严格只有一半非零权重的AST-Naive。这一点需要特别区分最接近稠密精度的结果使用了额外低秩参数而不是纯粹的3.4B参数2:4模型。十四、知识密集任务结果作者进一步测试了MMLU、MATH和GSM8K。方法WikiText困惑度MMLUMATHGSM8KLLaMA-2-7B稠密5.1245.35.3840.3Wanda 2:411.0227.62.8632.1AST-Naive5.8237.94.4235.6AST-Boosted5.6938.24.6436.2AST在困惑度上已经非常接近稠密模型但在MMLU上的差距仍然比较明显稠密模型45.3AST-Boosted38.2。这再次说明困惑度基本恢复不代表知识和推理能力已经完全恢复。不过相比Wanda的27.6AST显著保留了更多知识能力说明适量稀疏训练确实比一次性2:4剪枝更加可靠。GSM8K结果是在各模型都使用相同秩设置进行LoRA微调后获得的因此它反映的是剪枝模型的下游可适配能力而不是原始零样本数学能力。十五、消融实验说明了什么论文分别去除知识蒸馏、退火衰减和动态掩码。结果如下数值越低越好。方法GPT-2-124MGPT-2-350MGPT-2-774MOPT-125M普通稀疏训练40.3429.7928.6539.46不使用蒸馏39.2929.0827.2136.97固定衰减SR-STE32.8424.0421.7331.08固定掩码32.9324.1821.9531.06完整AST32.2323.6521.2930.22可以得到三个结论。第一知识蒸馏影响最大。不使用蒸馏时困惑度明显恶化说明有限训练数据下仅靠语言建模损失无法充分恢复稀疏模型。第二动态掩码确实有用。固定掩码始终略差于完整AST证明训练过程中重新选择连接能够修正初始剪枝错误。第三退火衰减优于固定衰减。二者差距不算巨大但退火策略能够在早期探索更多掩码并在后期获得更稳定的结构。十六、实际推理加速论文使用TensorRT-LLM测试LLaMA-2-7B 2:4稀疏模型的端到端吞吐量。RTX 4090输入长度输出长度稀疏吞吐量稠密吞吐量加速12812870.2352.941.33倍128102469.1152.001.33倍102412868.0651.101.33倍1024102467.4150.371.34倍NVIDIA L20输入长度输出长度稀疏吞吐量稠密吞吐量加速12812854.7529.861.83倍128102453.8129.571.82倍102412852.4929.181.80倍1024102451.6428.941.78倍虽然2:4删除了50%权重但实际端到端加速并不是理论上的2倍。原因包括并非所有模型算子都支持稀疏加速注意力Softmax、LayerNorm、RoPE等操作没有同步减少稀疏权重解码和调度存在开销不同GPU对2:4 Sparse Tensor Core的利用能力不同模型可能受到内存带宽或其他算子限制。L20上的收益明显高于RTX 4090也说明稀疏加速高度依赖具体硬件和软件栈。不过相比很多只报告FLOPs或理论稀疏率的论文AST至少给出了TensorRT-LLM上的实际端到端结果。十七、AST能否与量化结合论文还将AST与AWQ结合。在LLaMA-2-7B上方法理论存储比例WikiText困惑度稠密FP161.05.12AWQ 4-bit0.1255.68AST-Naive 8-bit0.1256.37AST-Naive 4-bit0.06756.48AST-Boosted 4-bit0.0786.25AST-Naive 4-bit的理论存储量约为原FP16模型的6.75%对应约14.8倍理论压缩。但这里的“理论存储比例”不能直接等同于实际模型文件大小或推理显存因为还需要考虑2:4掩码元数据量化缩放因子张量对齐未量化或未剪枝模块运行时工作空间。这部分实验主要证明2:4稀疏和低比特量化可以兼容并没有给出联合方案的完整端到端速度结果。十八、与上一篇Sparse Fine-tuning的区别对比维度Sparse Fine-tuningAST基础稀疏方式SparseGPT后固定掩码动态2:4幅值掩码掩码是否变化否是被剪权重能否恢复不能训练前中期可以核心蒸馏SquareHead中间特征蒸馏输出Logits的KL蒸馏是否约束隐藏层是实验认为不适合是否增加补偿参数无核心要求可选SLoRB主要硬件形式非结构化和规则N:M重点是2:4GPU端到端测速主要为局部内核TensorRT-LLM完整吞吐量Sparse Fine-tuning的核心观点是固定稀疏结构后通过逐层特征蒸馏恢复模型。AST的核心观点则是掩码本身也应该在重新训练过程中继续学习并通过退火机制从探索逐渐走向稳定。两篇论文对中间特征蒸馏给出了不同结论。这不是简单的对错关系而是说明蒸馏效果可能与任务、训练数据、稀疏类型和训练阶段密切相关。十九、与Prune and Tune的区别Prune and Tune通过多轮过程逐渐提高全局稀疏率剪到10% → 微调 → 剪到20% → 微调 → 继续增加AST则始终在2:4约束下训练每组始终只保留2个权重但保留的是哪两个可以不断变化。因此Prune and Tune调整的是整体稀疏率AST调整的是固定稀疏率下的连接位置。AST没有从低稀疏率逐渐增加到50%而是在训练期间保持2:4前向计算同时探索最合适的局部掩码。二十、方法优点20.1 同时优化权重和掩码AST没有把初始剪枝结果视为最终答案而是在训练过程中不断修正2:4连接位置。这比固定掩码恢复更适合严格的局部稀疏模式。20.2 探索和收敛之间设计合理退火衰减使模型前期可以尝试不同连接后期又能逐渐稳定避免掩码一直振荡。20.3 在知识任务上显著优于一次性剪枝AST虽然没有完全恢复MMLU和数学能力但相比Wanda 2:4已经明显改善说明重新训练对知识保持非常重要。20.4 使用的训练量远低于完整预训练LLaMA-2-7B只使用约7至7.5B Token远少于原始预训练语料。虽然成本仍然不低但比从头训练一个2:4模型现实得多。20.5 报告了真实GPU加速论文在RTX 4090和L20上使用TensorRT-LLM报告了1.33至1.83倍端到端吞吐量提升而不只是报告理论FLOPs。二十一、方法局限21.1 “少量训练成本”依然很高7至7.5B Token虽然不到原始预训练量的0.4%但仍然是一项昂贵的继续预训练任务。它需要多GPU训练稠密教师前向传播稀疏学生前向和反向传播AdamW优化器状态数十亿Token数据。因此AST远比Wanda和SparseGPT等一次性剪枝昂贵。21.2 大模型实验只有LLaMA-2-7B论文在GPT-2和OPT上测试了多个规模但真正达到数十亿参数的主要实验只有LLaMA-2-7B。没有验证LLaMA-2-13B或70BLLaMA-3Mistral更大的MoE模型长上下文模型。模型规模继续增大后动态掩码、教师蒸馏和优化器状态是否仍然经济并没有充分证明。21.3 最强结果依赖额外参数AST-Boosted加入12.5%的SLoRB参数后零样本平均准确率才达到58.62与稠密模型只差1.16个百分点。严格的AST-Naive只有57.68。因此论文最亮眼的精度结果并不是纯粹的50%参数删除而是“2:4稀疏主干加额外低秩分支”。21.4 训练阶段未必能够获得稀疏加速AST前向传播使用2:4掩码但训练仍需保存完整底层权重更新可能重新进入模型的权重保存优化器状态周期性重新排序生成掩码运行稠密教师。所以AST的主要硬件收益发生在最终推理阶段。论文没有证明整个重新训练过程也能获得与2:4推理相当的加速。21.5 知识能力仍没有完全恢复AST-Boosted的困惑度只比稠密模型高0.57但MMLU仍低7.1个百分点。这说明通用语言建模损失和Logit蒸馏并不能完整保护所有参数化知识与复杂推理能力。21.6 实际加速具有明显硬件依赖同一个2:4模型RTX 4090约加速1.33倍L20约加速1.8倍。换用不支持2:4 Sparse Tensor Core的GPU、普通PyTorch或其他推理框架可能无法获得同样收益。21.7 动态掩码仍然是局部幅值选择AST允许掩码变化但每组权重的选择仍然主要基于当前幅值。它没有直接考虑跨层误差全局任务敏感性Hessian信息不同层对延迟的实际贡献知识任务特定的重要性。因此“掩码可学习”并不意味着它执行了严格的全局可微掩码优化而是通过权重更新和周期性幅值重排序间接学习掩码。二十二、这篇论文真正有价值的地方AST最重要的贡献不是提出了一个新的静态剪枝分数而是改变了半结构化剪枝的处理方式。传统方法通常认为先找到掩码再恢复权重。AST则认为掩码和权重应该共同适应连接位置本身也是训练变量。它进一步说明严格2:4稀疏训练需要处理三个不同问题掩码探索允许错误删除被修正掩码收敛防止连接位置持续振荡知识恢复利用稠密教师避免稀疏模型只拟合有限训练数据。从部署角度看这篇论文也提供了一个比较完整的链路稠密模型 → 动态2:4稀疏训练 → 知识蒸馏 → 可选低秩补偿 → AWQ量化 → TensorRT-LLM推理它比只研究剪枝指标的一次性方法更接近实际可部署流程但代价是需要数十亿Token的继续训练。二十三、一句话总结AST在整个训练过程中保持2:4稀疏前向计算通过周期性幅值重排序允许被剪权重重新进入网络并使用由弱到强的退火衰减使掩码从早期探索逐渐过渡到后期稳定再结合稠密教师的Logit蒸馏和可选SLoRB低秩补偿LLaMA-2-7B的2:4模型可以将WikiText困惑度恢复到5.69并在TensorRT-LLM上获得约1.33至1.83倍GPU吞吐量提升但其最强结果依赖额外参数、数十亿Token重新训练和特定2:4硬件支持。

相关新闻