PyTorch 中网络剪枝、梯度剪裁、梯度累积

一、网络剪枝

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() 大模型训练、显存受限

三者可以同时使用:剪枝压缩模型体积,梯度累积解决显存问题,梯度剪裁保证训练稳定。

相关推荐
wiliam_luky2 小时前
RTMP和RTSP+RTP+RTCP
网络
zxanz12 小时前
SSL 证书有哪些类型,怎么选择适合自己使用的?
网络·网络协议·ssl
通信数码研究院2 小时前
2026户外监控选什么?TOP5场景适配榜单:五类安装环境与产品选型参考
网络
谢亮_vipxieliang2 小时前
容器日志收集与管理:从 stdout 规范到 ELK/Loki 落地
运维·网络·人工智能·elk·docker·容器
QYRdata2 小时前
年复合增长率19.6% 数据隐私安全软件赛道领跑新兴科技
网络·科技·服务发现
半仙白桑3 小时前
内核篇第十三讲:system-V消息队列
java·linux·网络
XUEYUAN52123 小时前
路由跳数与地理位置校验:为什么 IP 归属地和实际出口位置不一致
网络·网络协议·tcp/ip
Eloudy3 小时前
NVSwitch 和 UALink 的数据(缓存)一致性
网络·缓存·gpu
Multipath7123 小时前
Ku保底,Ka提速— 双频卫星聚合,打造应急通信“不断线”保底链路
网络·5g·安全·智能路由器·实时音视频