一、剪枝基础概念

  • 非结构化剪枝:随机剪去模型中 “不重要” 的单个权重参数(比如绝对值接近 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 目录")

总结

  1. 大模型剪枝核心分为非结构化剪枝(剪单个权重,易实现)和结构化剪枝(剪模型结构,硬件友好),Transformer 模型优先选注意力头剪枝。
  2. PyTorch 的torch.nn.utils.prune模块可实现通用剪枝,HuggingFace Transformers 为 Transformer 模型提供了专属的注意力头剪枝方法。
  3. 剪枝后需评估参数量、推理速度、精度三个维度,必要时通过微调恢复精度。

更多推荐