一、网络剪枝
PyTorch 提供了 torch.nn.utils.prune 模块,支持非结构化剪枝 (单个权重置零)和结构化剪枝(整通道/整滤波器移除)。
1. 非结构化剪枝(Magnitude-based)
按权重绝对值大小,移除最小的 k% 连接:
python
import torch
import torch.nn as nn
import torch.nn.utils.prune as prune
class SimpleNet(nn.Module):
def __init__(self):
super().__init__()
self.fc1 = nn.Linear(784, 256)
self.fc2 = nn.Linear(256, 10)
def forward(self, x):
x = torch.relu(self.fc1(x))
return self.fc2(x)
model = SimpleNet()
# 对 fc1 层的权重进行 L1 非结构化剪枝,剪掉 30% 的权重
prune.l1_unstructured(model.fc1, name="weight", amount=0.3)
# 查看剪枝后的权重(weight_orig 是原始权重,weight_mask 是掩码,weight 是剪枝后的结果)
print(f"剪枝后零值占比: {(model.fc1.weight == 0).sum().item() / model.fc1.weight.numel():.2%}")
# 永久固化剪枝结果(移除原始权重和掩码,只保留剪枝后的权重)
prune.remove(model.fc1, "weight")
2. 全局非结构化剪枝(跨层统一剪枝)
python
# 收集所有需要剪枝的层
parameters_to_prune = [
(model.fc1, "weight"),
(model.fc2, "weight"),
]
# 全局剪枝:所有层合在一起,剪掉绝对值最小的 40%
prune.global_unstructured(
parameters_to_prune,
pruning_method=prune.L1Unstructured,
amount=0.4,
)
# 固化
for module, name in parameters_to_prune:
prune.remove(module, name)
3. 结构化剪枝(通道剪枝)
结构化剪枝需要手动实现,因为 PyTorch 原生 prune 模块不直接支持通道级移除:
python
import torch
import torch.nn as nn
import copy
def channel_pruning(model, prune_ratio=0.3):
"""
基于 L1 范数的卷积层通道剪枝
:param model: 待剪枝模型
:param prune_ratio: 剪枝比例
:return: 剪枝后的新模型
"""
# Step 1: 收集所有 Conv2d 层的通道重要性(L1 范数)
importance_scores = []
conv_layers = []
for name, module in model.named_modules():
if isinstance(module, nn.Conv2d):
# 计算每个输出通道的 L1 范数: [out_channels]
channel_norm = module.weight.data.abs().mean(dim=(1, 2, 3))
importance_scores.append(channel_norm)
conv_layers.append((name, module))
# Step 2: 确定全局阈值
all_scores = torch.cat(importance_scores)
threshold = torch.quantile(all_scores, prune_ratio)
# Step 3: 为每个 Conv2d 层生成保留掩码
masks = {}
for name, module in conv_layers:
channel_norm = module.weight.data.abs().mean(dim=(1, 2, 3))
mask = (channel_norm > threshold).float() # 1=保留, 0=剪掉
masks[name] = mask
print(f"{name}: 保留 {mask.sum().item()}/{mask.numel()} 通道")
# Step 4: 构建新模型(需要手动重建,因为通道数变了)
# 这里简化处理:直接应用掩码到原模型(推理时跳过零通道)
for name, module in conv_layers:
mask = masks[name].view(-1, 1, 1, 1) # 广播到 [out_ch, in_ch, kH, kW]
module.weight.data *= mask
if module.bias is not None:
module.bias.data *= masks[name]
return model
# 使用示例
# pruned_model = channel_pruning(original_model, prune_ratio=0.3)
提示:结构化剪枝后通常需要微调恢复精度。剪枝比例建议从 10%~20% 开始,逐步增加。
二、梯度剪裁
梯度剪裁用于防止梯度爆炸,PyTorch 提供两种内置函数:
1. 按范数剪裁(推荐)
python
import torch
import torch.nn as nn
import torch.optim as optim
model = nn.Sequential(
nn.Linear(784, 256),
nn.ReLU(),
nn.Linear(256, 10)
)
optimizer = optim.Adam(model.parameters(), lr=1e-3)
criterion = nn.CrossEntropyLoss()
for batch_x, batch_y in dataloader:
optimizer.zero_grad()
output = model(batch_x)
loss = criterion(output, batch_y)
loss.backward()
# 按 L2 范数剪裁,max_norm 通常设为 0.5~5.0
torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=1.0)
optimizer.step()
2. 按值剪裁(逐元素限幅)
python
# 每个梯度元素限制在 [-clip_value, clip_value] 范围内
torch.nn.utils.clip_grad_value_(model.parameters(), clip_value=0.5)
3. 配合混合精度训练使用
使用 GradScaler 时,需要先 unscale_ 再剪裁:
python
from torch.cuda.amp import autocast, GradScaler
scaler = GradScaler()
for batch_x, batch_y in dataloader:
optimizer.zero_grad(set_to_none=True) # 更高效的清零方式
with autocast():
output = model(batch_x)
loss = criterion(output, batch_y)
scaler.scale(loss).backward()
# 先 unscale,再剪裁
scaler.unscale_(optimizer)
torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=1.0)
scaler.step(optimizer)
scaler.update()
max_norm 选择建议:RNN/Transformer 推荐 1.0~5.0,CNN 推荐 2.0~10.0。可以先不剪裁跑几个 batch,监控梯度范数分布,取 90 分位数作为阈值。
三、梯度累积
梯度累积用于在显存受限时模拟大批次训练。核心思想:多个小 batch 的梯度累加后,统一更新一次参数。
python
import torch
import torch.nn as nn
import torch.optim as optim
model = nn.Sequential(
nn.Linear(784, 256),
nn.ReLU(),
nn.Linear(256, 10)
).cuda()
optimizer = optim.AdamW(model.parameters(), lr=1e-3)
criterion = nn.CrossEntropyLoss()
# 梯度累积配置
accumulation_steps = 4 # 累积 4 个 batch 后更新一次
# 等效 batch_size = 实际 batch_size * accumulation_steps
for epoch in range(num_epochs):
for i, (batch_x, batch_y) in enumerate(dataloader):
batch_x, batch_y = batch_x.cuda(), batch_y.cuda()
# 前向传播 + 反向传播
output = model(batch_x)
loss = criterion(output, batch_y)
# 关键:损失除以累积步数,保证梯度尺度与大批次一致
loss = loss / accumulation_steps
loss.backward() # 梯度自动累加,不会覆盖
# 每 accumulation_steps 步更新一次参数
if (i + 1) % accumulation_steps == 0:
# 可选:梯度剪裁
torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=1.0)
optimizer.step()
optimizer.zero_grad(set_to_none=True) # 清零梯度
# 处理最后一个不完整的累积周期
if (i + 1) % accumulation_steps != 0:
optimizer.step()
optimizer.zero_grad(set_to_none=True)
关键注意事项
| 要点 | 说明 |
|---|---|
| 损失缩放 | loss / accumulation_steps 必须做,否则梯度会偏大 N 倍 |
| zero_grad 时机 | 只在 optimizer.step() 之后清零,中间步骤的梯度要保留 |
| BatchNorm 影响 | BN 统计量仍基于小 batch 计算,与真实大批次有差异 |
| 学习率调度 | 调度器步数应基于有效 batch 的更新次数,而非实际前向传播次数 |
| 本质 | 时间换空间,训练总时长不变甚至略增,但显存占用大幅降低 |
四、三者结合使用(完整训练循环)
python
import torch
import torch.nn as nn
import torch.optim as optim
import torch.nn.utils.prune as prune
from torch.cuda.amp import autocast, GradScaler
# 1. 定义模型
model = nn.Sequential(
nn.Linear(784, 256),
nn.ReLU(),
nn.Linear(256, 10)
).cuda()
# 2. 应用剪枝(训练前或训练中迭代剪枝)
prune.l1_unstructured(model[0], name="weight", amount=0.2)
optimizer = optim.AdamW(model.parameters(), lr=1e-3)
criterion = nn.CrossEntropyLoss()
scaler = GradScaler()
# 3. 训练配置
accumulation_steps = 4
max_grad_norm = 1.0
for epoch in range(num_epochs):
for i, (batch_x, batch_y) in enumerate(dataloader):
batch_x, batch_y = batch_x.cuda(), batch_y.cuda()
optimizer.zero_grad(set_to_none=True)
for accum_idx in range(accumulation_steps):
with autocast():
output = model(batch_x)
loss = criterion(output, batch_y) / accumulation_steps
scaler.scale(loss).backward()
# 4. 梯度剪裁 + 参数更新
scaler.unscale_(optimizer)
torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=max_grad_norm)
scaler.step(optimizer)
scaler.update()
# 5. 迭代剪枝:每轮 epoch 后重新剪枝并微调
# prune.l1_unstructured(model[0], name="weight", amount=0.05)
五、总结对比
| 技术 | 核心作用 | 关键 API | 适用场景 |
|---|---|---|---|
| 网络剪枝 | 减少参数量/计算量,模型压缩 | prune.l1_unstructured() / prune.global_unstructured() |
模型部署、边缘设备 |
| 梯度剪裁 | 防止梯度爆炸,稳定训练 | clip_grad_norm_() / clip_grad_value_() |
RNN、Transformer、深层网络 |
| 梯度累积 | 小显存模拟大批次训练 | loss / steps + 延迟 step() |
大模型训练、显存受限 |
三者可以同时使用:剪枝压缩模型体积,梯度累积解决显存问题,梯度剪裁保证训练稳定。