尧图网站设计 尧图网站设计YAOTU DESIGN
ARTICLE DETAIL

资讯详情

深耕网站设计与一线实操的经验洞察。

Wanda 权重激活剪枝实战指南:基于 AI-Research-SKILLs 实现 LLM 50% 稀疏化且精度损失小于 1%

Wanda 权重激活剪枝实战指南:基于 AI-Research-SKILLs 实现 LLM 50% 稀疏化且精度损失小于 1% AI 技能人工智能大模型深度学习【免费下载链接】AI-Research-SKILLsComprehensive open-source library of AI research and engineering skills for any AI model. Package the skills and your claude code/codex/gemini agent will be an AI research agent with full horsepower. Maintained by Orchestra Research.项目地址https://gitcode.com/gh_mirrors/ai/AI-Research-SKILLs点击查看免费下载本篇技术指南以 AI-Research-SKILLs 仓库中 Wanda 剪枝参考文档 为主体系统讲解 ICLR 2024 论文提出的WandaPruning by Weights and Activations算法如何仅凭权重幅值 × 输入激活范数这一简单准则在无需重训one-shot的前提下将 LLaMA 等大模型剪到 50% 稀疏度且精度损失低于 1%。读完本文你将掌握 Wanda 的核心剪枝准则、完整的 PyTorch 实现、校准数据构造、N:M 结构化稀疏扩展以及与 SparseGPT、幅度剪枝的选型对比并能在 model-pruning 技能 的指导下直接落地一条可复现的剪枝 → 评估 → 部署流水线。一、Wanda 是什么一种简单而有效的免重训剪枝方法WandaWeightand activation pruning源自 ICLR 2024 论文A Simple and Effective Pruning Approach for Large Language ModelsarXiv 编号 2306.11695官方实现位于locuslab/wanda仓库。它属于大语言模型一次性one-shot剪枝家族不需要重新训练模型只需一小段校准数据统计激活信息即可把模型权重剪到指定稀疏度。其最核心的卖点可以概括为一句话Wanda 用权重幅值 × 输入激活作为重要性度量在 50% 稀疏度下实现 1% 的精度损失且无需任何重训。在 AI-Research-SKILLs 仓库的体系里model-pruning 技能 归属于 19-emerging-techniques前沿技术类别与 MoE 训练、模型融合、长上下文、投机解码、知识蒸馏并列在 skill-routing 路由表 中它被明确路由用于Reducing model size缩小模型体积这一类研究任务。也就是说当研究 Agent 面临模型太大、推理太慢、硬件放不下的问题时就会调用该技能而 Wanda 正是该技能文档推荐的首选剪枝算法。Wanda 适用的典型场景压缩 40%60% 的模型体积精度损失控制在 1% 以内加速推理配合 2:4 / 4:8 这类 N:M 结构化稀疏模式在支持稀疏张量核心的硬件上获得 24 倍加速在移动端、边缘设备等受限硬件上部署降低显存与内存占用没有重训预算使用 one-shot 方法几分钟完成剪枝。二、核心创新重要性 权重幅值 × 激活使用度2.1 剪枝准则传统幅度剪枝magnitude pruning只考察权重自身的大小而 Wanda 的关键洞察是一个权重是否重要取决于它的幅值magnitude以及它被输入使用的频率usage。其重要性度量定义为importance(w_ij) |w_ij| × ||X_i||其中w_ij连接输入维度 i 到输出维度 j 的权重X_i输入维度 i 的激活向量||·||L2 范数实际实现中常用按输入维度的均值绝对范数近似。直觉解读权重幅值大 → 该参数本身承载的信息量大激活值高 → 该输入维度被频繁使用、对前向传播贡献大两者相乘同时捕获了静态重要性与动态使用度两层信号。2.2 与幅度剪枝的对比方法重要性度量特点幅度剪枝基线importance |weight|只考虑权重大小忽略使用频率Wandaimportance |weight| × activation同时考虑权重幅值与输入激活用一个具体例子说明差异Weight A: magnitude0.5, activation0.1 → importance0.05 Weight B: magnitude0.3, activation0.8 → importance0.24 幅度剪枝保留 A权重更大 Wanda 保留 B整体更重要✓幅度剪枝会保留幅值更大但几乎不被使用的权重 A而 Wanda 认为权重 B 虽然幅值略小却对应高频使用的输入维度综合重要性更高因此保留 B。这一准则上的差异正是 Wanda 在 50% 稀疏度下显著优于幅度剪枝平均精度仅损失 0.8% vs 4.9%的根本原因。三、算法实现一步步写出 Wanda 剪枝Wanda 的完整流程只有四步收集激活统计 → 计算重要性 → 按阈值生成掩码 → 应用剪枝。下面这份代码完整复现了 wanda.md 中给出的核心算法可直接运行import torch from transformers import AutoModelForCausalLM def wanda_prune(model, calib_data, sparsity0.5): Wanda pruning algorithm. Steps: 1. Collect activation statistics on calibration data 2. Compute importance |weight| × activation 3. Prune lowest importance weights 4. Return pruned model (no retraining!) # Step 1: Collect activations activations {} def activation_hook(name): def hook(module, input, output): # Store input activation norms X input[0].detach() # Per-input-dimension norm act_norm X.abs().mean(dim0) # Average over batch/sequence if name in activations: activations[name] act_norm else: activations[name] act_norm return hook # Register hooks hooks [] for name, module in model.named_modules(): if isinstance(module, torch.nn.Linear): hook module.register_forward_hook(activation_hook(name)) hooks.append(hook) # Run calibration model.eval() with torch.no_grad(): for batch in calib_data: model(**batch) # Remove hooks for hook in hooks: hook.remove() # Step 2 3: Prune based on importance for name, module in model.named_modules(): if isinstance(module, torch.nn.Linear) and name in activations: W module.weight.data act activations[name] # Compute importance (per output dimension) importance W.abs() * act.unsqueeze(0) # (out_features, in_features) # Find threshold for sparsity threshold torch.quantile(importance.flatten(), sparsity) # Create mask mask importance threshold # Apply pruning W.data * mask.float() return model几个值得注意的实现细节钩子hook只挂在torch.nn.Linear上LLM 中绝大部分参数量集中在线性层注意力投影、MLP 等Wanda 默认只剪这些层激活统计在eval()no_grad()下进行避免 dropout、BatchNorm 等训练期行为干扰统计同时省去梯度开销掩码直接就地乘到权重上W.data * mask.float()被剪掉的权重置 0不做任何重训或重建在 SKILL.md 的 Quick Start 版本中激活归一化用的是input[0].detach().abs().mean(dim0)即按批次/序列维度求均值与上面X.abs().mean(dim0)完全一致——两者的区别仅在于参考文档按累加多次校准前向的结果而 Quick Start 直接覆盖多次前向取累加能获得更稳定的统计量。3.1 Per-Output 剪枝按输出维度独立设定阈值Wanda 一个容易被忽略但至关重要的细节是剪枝是逐输出维度per-output进行的而非全局统一切一刀。# For each output dimension, prune sparsity% of weights for out_dim in range(out_features): # Importance for this output importance_out |W[out_dim, :]| × activation # Prune sparsity% of this outputs weights threshold quantile(importance_out, sparsity) mask_out importance_out threshold # Apply W[out_dim, :] * mask_out原因如果采用全局阈值某些输出行中权重普遍很大、另一些普遍很小就会导致部分输出被过度剪枝而失去表达能力。逐输出剪枝保证每个输出维度都保留相同的稀疏比例balanced pruning使各输出通道保持均等的容量从而显著降低整体精度损失。四、校准数据激活统计的来源Wanda 需要一小段校准数据来统计每层输入激活的分布。论文与官方实现的推荐配置如下参数推荐值说明样本量128 条论文实验中的默认规模数据来源任意文本语料C4、WikiText 等均可每条长度2048 tokens与常见 LLM 上下文长度对齐一个基于 HuggingFacedatasets的校准数据构造示例from datasets import load_dataset # Load calibration dataset calib_dataset load_dataset(allenai/c4, en, splittrain, streamingTrue) calib_samples [] for i, example in enumerate(calib_dataset): if i 128: break text example[text][:2048] # First 2048 chars calib_samples.append(text) # Tokenize tokenized tokenizer( calib_samples, return_tensorspt, paddingTrue, truncationTrue, max_length2048 )关于数据质量校准数据的质量越高剪枝效果会略好但并不关键参考文档明确标注not critical。这是因为 Wanda 只需要激活的近似分布来区分常用维度与冷门维度128 条样本已足够稳定地估计这一统计量。实际项目中完全可以用 SKILL.md 的 Quick Start 那样手写 3 条示例文本凑合验证流程再在生产环境换成 C4/WikiText 提升效果。五、性能结果来自 ICLR 2024 论文的实验证据以下结果均出自 Wanda 论文LLaMA 系列模型、零样本任务评测由 wanda.md 完整收录。5.1 非结构化稀疏Unstructured SparsityModelSparsityMethodPerplexity (WikiText2)Average AccuracyLLaMA-7B0%Baseline5.6860.2%LLaMA-7B50%Magnitude8.4555.3% (-4.9%)LLaMA-7B50%SparseGPT6.3259.1% (-1.1%)LLaMA-7B50%Wanda6.1859.4% (-0.8%)关键发现Wanda 在 50% 稀疏度下 perplexity6.18甚至优于依赖 Hessian 二阶信息的 SparseGPT6.32平均精度损失-0.8%也逼近 SparseGPT-1.1%却完全不需要复杂的矩阵求逆——以极简算法达到接近 SparseGPT 的质量正是论文标题中Simple and Effective的注脚。5.2 N:M 结构化稀疏硬件友好ModelSparsity PatternWanda PPLMagnitude PPLSpeedupLLaMA-7B2:4 (50%)6.429.122.0× (on A100)LLaMA-7B4:8 (50%)6.388.952.0× (on A100)N:M 稀疏每 M 个连续权重中保留 N 个与 NVIDIA 稀疏张量核心sparse tensor cores兼容可在 A100 等硬件上获得约 2 倍推理加速且 perplexity 相比非结构化方案仅小幅劣化6.42 vs 6.18远优于同模式的幅度剪枝9.12。5.3 向更大模型的扩展Model SizeSparsityWanda PPLDegradationLLaMA-7B50%6.180.50LLaMA-13B50%5.420.38LLaMA-30B50%4.770.21LLaMA-65B50%4.250.15扩展规律模型越大50% 剪枝带来的 perplexity 退化越小0.50 → 0.15。这印证了大规模 LLM 中存在大量冗余参数剪枝对它们的伤害更小——也意味着 Wanda 这类 one-shot 方法在大模型压缩场景中尤其有吸引力。六、扩展Wanda × N:M 结构化稀疏非结构化稀疏在没有专用硬件时无法直接带来推理加速Wanda 论文因此提供了 N:M 结构化变体。其思路与标准 Wanda 完全一致先用校准数据统计激活、计算重要性区别仅在于掩码生成方式不再是全局/逐输出阈值而是在每个长度为 M 的连续权重组内保留重要性最高的 N 个。def wanda_nm_prune(model, calib_data, n2, m4): Wanda with N:M structured sparsity. Keeps top-N weights per M consecutive weights. Compatible with NVIDIA sparse tensor cores. # Collect activations (same as standard Wanda) activations collect_activations(model, calib_data) # Prune with N:M pattern for name, module in model.named_modules(): if isinstance(module, torch.nn.Linear): W module.weight.data act activations[name] # Importance importance W.abs() * act.unsqueeze(0) # Apply N:M pruning W.data apply_nm_mask(W, importance, nn, mm) return model def apply_nm_mask(weight, importance, n2, m4): Apply N:M sparsity pattern. shape weight.shape # Flatten and pad to multiple of M importance_flat importance.flatten() weight_flat weight.flatten() pad_size (m - len(importance_flat) % m) % m importance_padded F.pad(importance_flat, (0, pad_size)) weight_padded F.pad(weight_flat, (0, pad_size)) # Reshape into groups of M importance_grouped importance_padded.reshape(-1, m) weight_grouped weight_padded.reshape(-1, m) # Find top-N per group _, indices torch.topk(importance_grouped, n, dim-1) # Create mask mask torch.zeros_like(importance_grouped) mask.scatter_(1, indices, 1.0) # Apply weight_pruned weight_grouped * mask weight_pruned weight_pruned.flatten()[:len(weight_flat)] return weight_pruned.reshape(shape)实现要点F.pad将扁平化后的权重补齐到 M 的整数倍避免最后一组长度不足torch.topk(..., n, dim-1)在每个长度为 m 的组内挑选重要性最高的 n 个索引scatter_将选中位置置 1生成 0/1 掩码剪枝后截断回原始长度并还原形状。在 SKILL.md 的 N:M 示例 中还给出了一个更简化的变体直接用weight_grouped.abs()做 topk纯幅度 N:M不依赖激活统计而 Wanda 版的 N:M 则始终基于importance |W| × act排序两者可视为幅度 N:M与激活感知 N:M两种策略后者的精度通常更优。七、Wanda vs SparseGPT如何选型SparseGPTarXiv 2301.00774是另一款知名的免重训剪枝方法它利用 Hessian 矩阵做逐层重建。两者的对比如下AspectWandaSparseGPTComplexityO(n) per layerO(n²) per layer (Hessian)SpeedFast (~minutes)Slow (~hours)MemoryLow (activations)High (Hessian matrix)Quality (50%)-0.8% accuracy-0.4% accuracyImplementationSimple (~100 lines)Complex (matrix inverse)权衡结论Wanda更简单、更快、显存占用低50% 稀疏度下精度损失约 0.8%SparseGPT更复杂、更慢需要 Hessian 逆但精度损失可低至 0.4%推荐除非你的场景对精度有极致要求否则优先使用 Wanda。SKILL.md 的方法选择逻辑 也给出了同样的判断无重训预算选 Wanda更快追求极致质量选 SparseGPT追求硬件加速选 N:M。八、实践部署完整剪枝脚本与评估8.1 使用官方 CLI 一键剪枝参考文档提供了基于官方locuslab/wanda仓库的完整命令行流程# Clone Wanda repo git clone https://github.com/locuslab/wanda cd wanda # Install dependencies pip install torch transformers datasets # Prune LLaMA-7B to 50% sparsity python main.py \ --model meta-llama/Llama-2-7b-hf \ --prune_method wanda \ --sparsity_ratio 0.5 \ --sparsity_type unstructured \ --save ./pruned_models/llama-7b-wanda-50 # Prune with 2:4 structured sparsity (NVIDIA GPUs) python main.py \ --model meta-llama/Llama-2-7b-hf \ --prune_method wanda \ --sparsity_ratio 0.5 \ --sparsity_type 2:4 \ --save ./pruned_models/llama-7b-wanda-2-4关键参数说明参数含义取值建议--model待剪枝的 HF 模型名或路径LLaMA、Mistral 等因果语言模型--prune_method剪枝算法wanda、sparsegpt、magnitude等--sparsity_ratio稀疏比例0.5 表示剪掉 50% 权重--sparsity_type稀疏模式unstructured或2:4、4:8等 N:M 模式--save剪枝后模型输出目录可用save_pretrained继续导出其中--sparsity_type 2:4会调用第六节的 N:M 掩码逻辑产出与 NVIDIA sparse tensor cores 兼容的模型。8.2 用 lm-evaluation-harness 评估剪枝效果剪枝完成后用标准评测库对比原模型与剪枝模型的零样本准确率from lm_eval import evaluator # Evaluate pruned model results evaluator.simple_evaluate( modelhf, model_argspretrained./pruned_models/llama-7b-wanda-50, tasks[arc_easy, arc_challenge, hellaswag, winogrande], batch_size8 ) print(Accuracy after 50% pruning:) for task, score in results[results].items(): print(f{task}: {score[acc]:.3f})评测建议参照 SKILL.md 的评估章节同时评测原模型与剪枝模型逐一对比各任务 acc 与 perplexity把退化量控制在可接受范围50% 稀疏度下期望 Wanda 1%、SparseGPT 0.5%、幅度剪枝 23%。注意这些数值来自 Wanda / SparseGPT 论文在 LLaMA-7B 上的报告不同模型与数据集上会有浮动应以本地实测为准。8.3 生产流水线剪枝 可选微调恢复若精度仍不满足可在剪枝后追加轻量微调参考 SKILL.md 的 production_pruning_pipeline加载模型 → 用 C4/WikiText 前 1000 条做校准 → 执行 Wanda/SparseGPT 剪枝 → 用Trainer以 1e-5 学习率微调 1 个 epoch 恢复精度 →save_pretrained导出。此外 SKILL.md 还提供了三种进阶剪枝策略渐进式幅度剪枝训练过程中按步数从 0% 线性提升到目标稀疏度逐层差异化剪枝早层少剪如 30%、晚层多剪如 60%因为早期层承载更多语义信息迭代剪枝 微调每轮小幅提升稀疏度后微调 2 个 epoch在高稀疏度70%下明显优于一次性剪枝。九、局限性使用前必须知道的边界不可重训恢复Wanda 是 one-shot 方法一旦剪错无法通过自身机制挽回若稀疏度过高导致严重退化只能依赖迭代剪枝或外部微调补偿依赖校准数据激活统计完全来自校准集校准数据与真实部署分布的偏移会影响剪枝质量尽管影响有限非结构化稀疏无加速不借助 N:M 等结构化模式时稀疏权重矩阵在普通硬件上无法直接提速甚至可能更慢要获得推理加速必须搭配支持稀疏计算的硬件如 NVIDIA sparse tensor cores或配套稀疏内核。十、在本仓库中的定位与延伸阅读技能入口model-pruning/SKILL.md 提供从安装、Quick Start、SparseGPT/N:M 变体到生产部署的完整指南是本篇参考文档的配套实操层参考文档references/wanda.md 即本文核心素材聚焦 Wanda 算法本身的原理、代码与实验数据路由上下文skill-routing.md 将该技能映射到缩小模型体积Reducing model size任务说明它是研究 Agent 在处理模型压缩需求时被路由到的目标技能相邻技术同属压缩方向的 知识蒸馏teacher-student 压缩可作为剪枝之外的另一条模型瘦身路径二者常被结合使用。结合本文内容一条典型的 Agent 工作流是研究任务提出在 1×A100 上以最小精度损失部署 7B 模型 → 路由到 model-pruning 技能 → 按本文流程用 C4 校准数据执行 Wanda 50% 剪枝 → 用 lm-evaluation-harness 对比精度 → 若需推理加速则改用 2:4 结构化稀疏。整个过程几分钟即可完成无需任何重训预算。赞分享AI 技能人工智能大模型深度学习【免费下载链接】AI-Research-SKILLsComprehensive open-source library of AI research and engineering skills for any AI model. Package the skills and your claude code/codex/gemini agent will be an AI research agent with full horsepower. Maintained by Orchestra Research.项目地址https://gitcode.com/gh_mirrors/ai/AI-Research-SKILLs点击查看免费下载相关推荐深度神经网络稀疏化实战指南基于 state_of_sparsity 复现 Transformer 与 ResNet-50 的剪枝与稀疏训练深度神经网络稀疏化实战指南基于 state_of_sparsity 复现 Transformer 与 ResNet 50 的剪枝与稀疏训练 导读 本指南以 s人工智能深度学习NLP计算机视觉强化学习YOLOv5 模型剪枝与稀疏度实战指南val.py 基线测试、0.3 稀疏度剪枝与精度影响解析YOLOv5 模型剪枝与稀疏度实战指南val.py 基线测试、0.3 稀疏度剪枝与精度影响解析 模型剪枝Pruning是提升深度模型推理效率的核心手段之一人工智能深度学习计算机视觉5分钟上手OpenRCT2插件开发用JavaScript为游乐园写第一个自定义窗口5分钟上手OpenRCT2插件开发用JavaScript为游乐园写第一个自定义窗口 OpenRCT2 是《过山车大亨2》的开源重制版它内置了一套完整的游戏开发图形学创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考
返回列表