大模型剪枝代码实战|从零到一,直接跑通!
·
一、剪枝基础概念
- 非结构化剪枝:随机剪去模型中 “不重要” 的单个权重参数(比如绝对值接近 0 的权重),实现简单但硬件加速效果有限。
- 结构化剪枝:剪去模型的完整结构(比如 Transformer 的注意力头、整个卷积层),硬件友好,是工业界主流。
- 核心指标:稀疏度 = 被剪枝的参数数量 / 总参数数量(比如稀疏度 50% 表示剪去一半参数)。
二、环境准备
首先安装依赖库(建议用虚拟环境)
# 核心依赖:PyTorch + Transformers(加载预训练模型)+ accelerate(加速推理)
pip install torch transformers accelerate datasets
选择轻量级的bert-base-uncased作为演示模型(约 110M 参数),先实现非结构化权重剪枝,再实现 Transformer 特有的注意力头剪枝,最后评估剪枝效果。
1. 完整可运行代码
import torch
import torch.nn.utils.prune as prune
from transformers import BertForSequenceClassification, BertTokenizer, pipeline
from time import time
import numpy as np
# ===================== 1. 基础配置 =====================
# 设备配置:优先用GPU,没有则用CPU
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
# 预训练模型名称(小型BERT,适合快速测试)
model_name = "bert-base-uncased"
# 剪枝配置
sparsity = 0.5 # 非结构化剪枝稀疏度(剪去50%的权重)
prune_heads_ratio = 0.2 # 注意力头剪枝比例(剪去20%的头)
# ===================== 2. 加载模型和Tokenizer =====================
# 加载BERT分类模型(用于文本分类任务)
model = BertForSequenceClassification.from_pretrained(
model_name,
num_labels=2, # 二分类任务(比如情感分析)
torch_dtype=torch.float32
).to(device)
# 加载Tokenizer
tokenizer = BertTokenizer.from_pretrained(model_name)
# 定义评估函数:计算模型参数量和推理速度
def evaluate_model(model, tokenizer, device):
"""评估模型:返回参数量、推理耗时"""
# 1. 计算参数量
total_params = sum(p.numel() for p in model.parameters())
trainable_params = sum(p.numel() for p in model.parameters() if p.requires_grad)
# 2. 测试推理速度
classifier = pipeline(
"text-classification",
model=model,
tokenizer=tokenizer,
device=device.index if torch.cuda.is_available() else -1
)
test_text = "I love programming with large language models!"
start_time = time()
for _ in range(100):
classifier(test_text)
avg_time = (time() - start_time) / 100 # 平均推理耗时(秒/次)
return {
"total_params": total_params / 1e6, # 转换为百万
"trainable_params": trainable_params / 1e6,
"avg_infer_time": avg_time
}
# ===================== 3. 原始模型评估 =====================
print("===== 原始模型评估 =====")
original_metrics = evaluate_model(model, tokenizer, device)
print(f"总参数量:{original_metrics['total_params']:.2f}M")
print(f"可训练参数量:{original_metrics['trainable_params']:.2f}M")
print(f"平均推理耗时:{original_metrics['avg_infer_time']:.4f}秒/次")
# ===================== 4. 非结构化权重剪枝(核心步骤) =====================
def prune_bert_unstructured(model, sparsity, device):
"""对BERT的全连接层进行非结构化剪枝"""
# 遍历模型的每一层,对全连接层(Linear)的weight进行剪枝
for name, module in model.named_modules():
# 只剪Linear层的权重(BERT的核心计算层)
if isinstance(module, torch.nn.Linear):
# 使用L1范数剪枝:移除绝对值最小的权重(PyTorch内置方法)
prune.l1_unstructured(
module,
name="weight", # 剪枝权重参数
amount=sparsity # 剪枝比例
)
# 永久化剪枝(移除剪枝掩码,真正删除参数)
prune.remove(module, "weight")
return model
# 执行非结构化剪枝
print("\n===== 执行非结构化剪枝(稀疏度50%) =====")
model_pruned = prune_bert_unstructured(model, sparsity, device)
# ===================== 5. Transformer注意力头剪枝(结构化剪枝) =====================
def prune_bert_heads(model, prune_ratio, device):
"""对BERT的注意力头进行结构化剪枝"""
# 遍历所有Transformer层
for layer_idx, layer in enumerate(model.bert.encoder.layer):
# 获取注意力头数量
num_heads = layer.attention.self.num_attention_heads
# 计算要剪的头数量
num_prune_heads = int(num_heads * prune_ratio)
# 随机选择要剪的头(也可以按重要性排序剪)
prune_heads = np.random.choice(num_heads, num_prune_heads, replace=False).tolist()
# 执行头剪枝(HuggingFace Transformers内置方法)
layer.attention.prune_heads(prune_heads)
return model
# 执行注意力头剪枝
print("\n===== 执行注意力头剪枝(比例20%) =====")
model_pruned = prune_bert_heads(model_pruned, prune_heads_ratio, device)
# ===================== 6. 剪枝后模型评估 =====================
print("\n===== 剪枝后模型评估 =====")
pruned_metrics = evaluate_model(model_pruned, tokenizer, device)
print(f"总参数量:{pruned_metrics['total_params']:.2f}M(减少{original_metrics['total_params'] - pruned_metrics['total_params']:.2f}M)")
print(f"可训练参数量:{pruned_metrics['trainable_params']:.2f}M")
print(f"平均推理耗时:{pruned_metrics['avg_infer_time']:.4f}秒/次(提速{(original_metrics['avg_infer_time'] - pruned_metrics['avg_infer_time'])/original_metrics['avg_infer_time']*100:.2f}%)")
# ===================== 7. 保存剪枝后的模型 =====================
print("\n===== 保存剪枝后模型 =====")
model_pruned.save_pretrained("./bert_pruned")
tokenizer.save_pretrained("./bert_pruned")
print("剪枝后模型已保存到 ./bert_pruned 目录")
总结
- 大模型剪枝核心分为非结构化剪枝(剪单个权重,易实现)和结构化剪枝(剪模型结构,硬件友好),Transformer 模型优先选注意力头剪枝。
- PyTorch 的
torch.nn.utils.prune模块可实现通用剪枝,HuggingFace Transformers 为 Transformer 模型提供了专属的注意力头剪枝方法。 - 剪枝后需评估参数量、推理速度、精度三个维度,必要时通过微调恢复精度。
更多推荐

所有评论(0)