模型剪枝经典论文精读:Towards Efficient Model Compression via Learned Global Ranking

一、论文基本信息

论文题目:Towards Efficient Model Compression via Learned Global Ranking

常用简称:LeGR

早期题目/常见引用名:LeGR: Filter Pruning via Learned Global Ranking

作者:Ting-Wu Chin、Ruizhou Ding、Cha Zhang、Diana Marculescu

发表信息:CVPR 2020 Oral

论文链接:https://arxiv.org/abs/1904.12368

官方代码:https://arxiv.org/abs/1904.12368

这篇论文发表于 CVPR 2020,CVF 页面显示论文题目为 Towards Efficient Model Compression via Learned Global Ranking,收录于 Proceedings of the IEEE/CVF Conference on Computer Vision and Pattern Recognition (CVPR), 2020, pp. 1518–1528。arXiv 页面标注该论文为 CVPR 2020 Oral,代码仓库说明其为 CVPR 2020 Oral paper 的代码实现。

LeGR 的核心目标不是只生成一个满足固定 FLOPs 约束的剪枝模型,而是希望 学习一次全局 filter 排名,然后通过剪掉排名靠后的 filters,快速得到一组不同精度—速度权衡的子网络。论文摘要明确指出,LeGR 学习跨层 filter 的全局排名,用来获得多个具有不同 accuracy/latency trade-off 的 ConvNet 架构。


二、论文要解决的问题

在 LeGR 之前,大多数结构化 filter pruning 方法都需要用户先指定一个目标复杂度,例如:

目标 FLOPs
目标参数量
目标模型大小
目标推理延迟

然后剪枝算法围绕这个固定目标生成一个剪枝模型。

这个流程的问题是:真实应用中,我们往往很难一开始就知道哪个复杂度最合适。比如机器人、无人机、移动端应用中,模型速度和模型精度都会影响最终系统表现;而最终任务效果往往只有在部署推理后才能真正判断。论文指出,寻找 accuracy 和 speed 的 sweet spot 往往需要反复试错,因此只针对一个预定义复杂度生成单个模型会很低效。

传统流程可以概括为:

指定一个 FLOPs 目标
    ↓
运行一次剪枝算法
    ↓
fine-tune
    ↓
评估精度和速度
    ↓
不满意,再换一个 FLOPs 目标重来

LeGR 想解决的问题是:

能不能只学习一次剪枝排序,然后快速得到多个不同 FLOPs / latency / accuracy 权衡的模型?

也就是说,LeGR 把目标从:

生成一个固定复杂度的剪枝模型

改成:

生成一组不同复杂度的剪枝模型,
方便用户选择精度和速度之间的最佳折中点。

论文明确提出,要把剪枝目标从输出单个预定义复杂度的 ConvNet,改成输出一组具有不同 accuracy/speed trade-off 的 ConvNets。


三、核心思想

LeGR 的核心思想可以概括为一句话:

先学习一个跨层的全局 filter 排名,然后根据不同 FLOPs 目标,从低排名到高排名依次删除 filters,从而一次排序生成多个剪枝模型。

这和传统 L1 / L2 filter pruning 不同。

传统 L1 / L2 剪枝通常是在每一层内部排序:

conv1 内部排一次
conv2 内部排一次
conv3 内部排一次
……

这种方式只能知道同一层内部哪些 filters 相对不重要,但不能直接比较:

conv2 的第 5 个 filter
和
conv8 的第 20 个 filter
到底谁更应该先被剪?

LeGR 的关键就是把不同层的 filters 放到同一个全局排序表里:

global ranking:
    filter_1
    filter_2
    filter_3
    ...
    filter_K

然后,对于任意一个 FLOPs 目标,只需要从排名最低的 filter 开始删除,直到满足目标复杂度即可。

这意味着:

学习一次 ranking
    ↓
剪少一点,得到高精度模型
    ↓
剪多一点,得到高加速模型
    ↓
继续剪,得到更激进压缩模型

论文中明确说,LeGR 的 global ranking 只需要为一个 ConvNet 学习一次,然后可以用来得到不同 FLOP count 的 ConvNets。


四、方法细节

4.1 为什么不能直接全局比较 L2 范数?

很多 filter pruning 方法使用 filter weight 的 L1 或 L2 范数作为重要性指标。LeGR 也承认 filter norm 在同一层内部是有意义的,但它指出一个关键问题:

filter norm 可以用于同层内部比较,
但不能直接跨层比较。

原因是不同层的 filter 尺寸、输入通道数、输出分布、BN 影响、残差连接位置都不同。某一层中的小范数 filter 不一定比另一层中的大范数 filter 更不重要。

论文把这一点写成 Norm Assumption

L2 norm 可以用于比较同一层内部 filters 的重要性,
但不能直接用于跨层比较。

论文随后提出:为了让 filter norms 能跨层比较,需要学习每一层自己的仿射变换。


4.2 LeGR 的 filter 重要性公式

LeGR 对第 (i) 个 filter 的重要性定义为:I_i =\alpha_{l(i)} \left| \Theta_i \right|2^2 + \kappa{l(i)}
 

其中:

  • \Theta_i:第 (i) 个 filter 的权重;

  • l(i):第 (i) 个 filter 所在层的编号;

  • |\Theta_i|_2^2:该 filter 的 L2 范数平方;

  • \alpha_{l(i)}:该层的可学习缩放系数;

  • \kappa_{l(i)}:该层的可学习平移系数。

这个公式是 LeGR 的核心。它不是直接用原始 filter norm 排序,而是对每一层的 norm 做一个 layer-wise affine transformation

原始 filter norm
    ↓
每层单独缩放和平移
    ↓
变成可跨层比较的 global importance

论文明确写道,LeGR 通过学习 layer-wise scale 和 shift 参数\alpha,\kappa ,把各层 filter norms 映射到一个全局可比较的重要性空间。


4.3 为什么只学习每层的\alpha  和\kappa

假设网络一共有 (K) 个 filters。如果我们直接学习每个 filter 的任意全局排序,搜索空间会非常巨大。论文指出,获得最优 global ranking 可能需要O(K \times K!) 轮 fine-tuning,几乎不可行。

所以 LeGR 做了一个折中:

同层内部:
    仍然相信 L2 norm 排序。

跨层之间:
    学习每层的缩放和平移,让不同层可比较。

这样,全局排序不再需要为每个 filter 学习独立参数,而只需要为每一层学习两个参数:\alpha_l,\ \kappa_l

如果网络有 (L) 个可剪卷积层,那么只需要学习 (2L) 个参数。

这使得问题从:

直接搜索 K 个 filters 的全排列

变成:

搜索 L 层的 affine transformation 参数

大大降低了搜索难度。


4.4 LeGR-Pruning 如何得到一个子网络?

给定已经学好的 \alpha,\kappa,LeGR 的剪枝过程很简单:

1. 计算每个 filter 的 L2 norm。
2. 用对应层的 alpha 和 kappa 做仿射变换。
3. 得到每个 filter 的全局重要性 I_i。
4. 按 I_i 从小到大排序。
5. 从排名最低的 filters 开始删除。
6. 一直删到满足目标 FLOPs。
7. 得到剪枝结构。
8. fine-tuning 剪枝模型。

论文 Figure 2 描述的就是这个流程:给定学习好的  \alpha,\kappa,LeGR-Pruning 返回 filter masks;剪枝后再 fine-tune 得到最终模型。

这一步有一个很大的优点:一旦 \alpha,\kappa 学好,后续生成不同 FLOPs 的模型不再需要重新搜索。对于多个 FLOPs 目标,只需要改变“剪到什么位置”为止。


4.5 关键假设:Subset Assumption

LeGR 之所以能通过一个 global ranking 得到多个不同 FLOPs 的模型,是因为它隐含了一个重要假设:更小模型的 filters 是更大模型 filters 的子集。

论文称之为 Subset Assumption

直观理解:

如果 80% FLOPs 模型保留 filters 集合 A,
60% FLOPs 模型保留 filters 集合 B,
那么 B 应该是 A 的子集。

也就是说,越小的模型只是从大模型中继续删掉更多低排名 filters。

论文承认这个假设比较强,因为实际最优结构未必严格满足嵌套关系:某个更小模型可能在某些层保留更多 filters,而在另一些层保留更少 filters。但这个假设能让一次 global ranking 生成一整条 trade-off curve,显著提高效率;实验也表明,在这个假设下得到的剪枝模型仍具有竞争力。


4.6 如何学习 \alpha 和 \kappa

LeGR 把 \alpha,\kappa 的学习看成一个超参数优化问题。

目标是:\arg\max_{\alpha,\kappa} \operatorname{Acc}{val} \left( \hat{\Theta}{l} \right)

其中:\hat{\Theta}_{l} = \operatorname{LeGR\text{-}Pruning} (\alpha,\kappa,\hat{\zeta}_l)

这里  \hat{\zeta}_l表示最低考虑的 FLOPs 约束,也就是搜索时用的最小模型复杂度。论文的做法是:用最低 FLOPs 约束生成剪枝模型,短暂 fine-tuning 后在验证集上评估 accuracy,把这个 accuracy 作为当前 \alpha,\kappa的 fitness。

为什么只看最低 FLOPs 约束?

因为不同 FLOPs 下模型准确率天然不同,很难直接比较。LeGR 选择在最激进约束下学习 ranking,认为如果一个 global ranking 在低 FLOPs 模型上仍然有效,那么它通常也能为更高 FLOPs 模型提供有用排序。论文后续也做了 \hat{\zeta}_l的消融。


4.7 Regularized Evolutionary Algorithm

LeGR 使用 Regularized Evolutionary Algorithm,EA 来搜索\hat{\zeta}_l

基本过程是:

1. 初始化一个候选池 Pool。
2. 每个候选是一个 alpha-kappa pair。
3. 对候选执行 LeGR-Pruning。
4. 对剪枝模型短暂 fine-tuning。
5. 用验证集 accuracy 作为 fitness。
6. 从候选池中采样一组候选。
7. 选择 fitness 最好的候选。
8. 随机选择一部分层,对这些层的 alpha/kappa 做 mutation。
9. 生成新候选。
10. 用新候选替换候选池中最老的候选。
11. 重复搜索。

论文 Algorithm 1 给出了这个流程。它使用 regularized EA,随机选择一部分层进行 mutation,并用验证集精度作为 fitness。论文还提到,搜索阶段使用短 fine-tuning 近似完整 fine-tuning,例如实验中使用 (\hat{\tau}=200) 个 gradient updates。

这和 EagleEye 有点相似:两者都不想对每个候选完整训练到收敛。

区别是:

EagleEye:
    重点是快速评估候选子网;
    用 Adaptive BN 提高评估可靠性。

LeGR:
    重点是学习全局 filter ranking;
    用短 fine-tuning + EA 学习 alpha-kappa。

4.8 依赖通道如何处理?

真实网络中并不是每个 filter 都能独立删除。比如:

Depthwise convolution 中的通道
Residual connection 中相加的通道
分支结构中需要对齐的通道

这些通道之间存在依赖关系。如果只删其中一个,可能导致维度不匹配。

论文说明,LeGR 对有依赖的 channels 进行分组,并联合剪枝。具体来说,对于 depth-wise convolution,会把其通道和前一层对应通道分组;对于 residual connection 中相加的 channels,也会分组一起剪。重要性度量则使用 learned affine transformation 后的重要性。

这说明 LeGR 不只是理论排序,也考虑了实际 CNN 结构中的通道依赖问题。

"""
LeGR Toy Demo: Learned Global Ranking for 5-layer CNN
=====================================================

对应论文:
    Towards Efficient Model Compression via Learned Global Ranking
    Ting-Wu Chin et al., CVPR 2020

这份代码是 LeGR 的“理解性代码”,不是官方代码逐行复现。

核心演示:
    1. 构建一个 5 层 CNN。
    2. 支持随机数据 / CIFAR-10。
    3. 计算每个卷积 filter 的 L2 norm。
    4. 用每层的 alpha / kappa 做仿射变换:

           I_i = alpha_l * ||Theta_i||_2^2 + kappa_l

    5. 把所有层的 filters 放到同一个全局排序表里。
    6. 对不同 global prune ratio,从全局低重要性 filters 开始删除。
    7. 生成多个剪枝方案和真实窄网络。
    8. 统计参数量和 MACs。
    9. 可选:用一个简化 evolutionary search 学习 alpha/kappa。

运行示例:

    # 只演示 LeGR 全局排序和多预算剪枝,不训练,最快
    python legr_toy_5conv_global_ranking_demo.py --data random

    # 使用 CIFAR-10,训练 baseline 1 epoch 后再做排序和剪枝
    python legr_toy_5conv_global_ranking_demo.py --data cifar10 --epochs 1

    # 用手动 alpha/kappa 演示跨层重标定
    python legr_toy_5conv_global_ranking_demo.py \
        --alpha 1.0 0.8 1.2 1.5 0.6 \
        --kappa 0.0 0.1 0.0 -0.1 0.2

    # 生成 10%、30%、50% 三个剪枝预算下的窄模型
    python legr_toy_5conv_global_ranking_demo.py \
        --global-prune-ratios 0.1 0.3 0.5

    # 简化版搜索 alpha/kappa:为了理解流程,不追求论文级结果
    python legr_toy_5conv_global_ranking_demo.py \
        --data cifar10 --epochs 1 --search-alpha-kappa --search-iters 20 --search-target-ratio 0.4

注意:
    - LeGR 原论文用 regularized evolutionary algorithm 学 alpha/kappa。
    - 本 demo 的搜索部分是简化版,主要用于理解“学习跨层仿射变换”的思想。
    - 真实论文实验需要大模型、真实训练策略、短 fine-tuning、完整验证集。
"""

from __future__ import annotations

import argparse
import copy
import math
import random
from dataclasses import dataclass
from typing import Dict, Iterable, List, Optional, Sequence, Tuple

import torch
import torch.nn as nn
import torch.nn.functional as F
import torch.utils.data as data

try:
    from torchvision import datasets, transforms
except Exception:
    datasets = None
    transforms = None


# ============================================================
# 1. 5 层 CNN
# ============================================================

class Tiny5ConvNet(nn.Module):
    """
    一个便于物理剪枝的 5 层 CNN。

    默认通道:
        conv1: 3  -> 8
        conv2: 8  -> 12
        conv3: 12 -> 16
        conv4: 16 -> 16
        conv5: 16 -> 10

    每层:Conv -> BN -> ReLU。
    在 conv1、conv2、conv4 后做 MaxPool。
    最后 AdaptiveAvgPool + Linear。
    """

    def __init__(self, channels: Sequence[int] = (8, 12, 16, 16, 10), num_classes: int = 10) -> None:
        super().__init__()
        if len(channels) != 5:
            raise ValueError("Tiny5ConvNet expects exactly 5 conv channel numbers.")

        self.channels = list(map(int, channels))
        self.num_classes = int(num_classes)

        in_channels = [3] + self.channels[:-1]
        out_channels = self.channels

        self.convs = nn.ModuleList()
        self.bns = nn.ModuleList()

        for cin, cout in zip(in_channels, out_channels):
            self.convs.append(nn.Conv2d(cin, cout, kernel_size=3, padding=1, bias=False))
            self.bns.append(nn.BatchNorm2d(cout))

        self.pool = nn.AdaptiveAvgPool2d((1, 1))
        self.fc = nn.Linear(self.channels[-1], self.num_classes)

        self._init_weights()

    def _init_weights(self) -> None:
        for m in self.modules():
            if isinstance(m, nn.Conv2d):
                nn.init.kaiming_normal_(m.weight, mode="fan_out", nonlinearity="relu")
            elif isinstance(m, nn.BatchNorm2d):
                nn.init.ones_(m.weight)
                nn.init.zeros_(m.bias)
            elif isinstance(m, nn.Linear):
                nn.init.normal_(m.weight, mean=0.0, std=0.01)
                nn.init.zeros_(m.bias)

    def forward(self, x: torch.Tensor) -> torch.Tensor:
        for i in range(5):
            x = self.convs[i](x)
            x = self.bns[i](x)
            x = F.relu(x, inplace=True)

            if i in (0, 1, 3):
                x = F.max_pool2d(x, kernel_size=2, stride=2)

        x = self.pool(x).flatten(1)
        x = self.fc(x)
        return x


# ============================================================
# 2. 数据加载
# ============================================================

def set_seed(seed: int) -> None:
    random.seed(seed)
    torch.manual_seed(seed)
    torch.cuda.manual_seed_all(seed)


def build_random_loaders(
    num_train: int,
    num_val: int,
    batch_size: int,
    num_classes: int = 10,
) -> Tuple[data.DataLoader, data.DataLoader]:
    train_images = torch.randn(num_train, 3, 32, 32)
    train_labels = torch.randint(0, num_classes, (num_train,))
    val_images = torch.randn(num_val, 3, 32, 32)
    val_labels = torch.randint(0, num_classes, (num_val,))

    train_loader = data.DataLoader(
        data.TensorDataset(train_images, train_labels),
        batch_size=batch_size,
        shuffle=True,
        num_workers=0,
    )
    val_loader = data.DataLoader(
        data.TensorDataset(val_images, val_labels),
        batch_size=batch_size,
        shuffle=False,
        num_workers=0,
    )
    return train_loader, val_loader


def build_cifar10_loaders(
    data_root: str,
    batch_size: int,
    num_workers: int,
    train_subset: int,
    val_subset: int,
) -> Tuple[data.DataLoader, data.DataLoader]:
    if datasets is None or transforms is None:
        raise ImportError("torchvision is required for CIFAR-10 mode.")

    mean = (0.4914, 0.4822, 0.4465)
    std = (0.2470, 0.2435, 0.2616)

    train_tf = transforms.Compose([
        transforms.RandomCrop(32, padding=4),
        transforms.RandomHorizontalFlip(),
        transforms.ToTensor(),
        transforms.Normalize(mean, std),
    ])
    test_tf = transforms.Compose([
        transforms.ToTensor(),
        transforms.Normalize(mean, std),
    ])

    train_set = datasets.CIFAR10(data_root, train=True, download=True, transform=train_tf)
    val_set = datasets.CIFAR10(data_root, train=False, download=True, transform=test_tf)

    if train_subset > 0:
        train_set = data.Subset(train_set, list(range(min(train_subset, len(train_set)))))
    if val_subset > 0:
        val_set = data.Subset(val_set, list(range(min(val_subset, len(val_set)))))

    train_loader = data.DataLoader(
        train_set,
        batch_size=batch_size,
        shuffle=True,
        num_workers=num_workers,
        pin_memory=True,
    )
    val_loader = data.DataLoader(
        val_set,
        batch_size=batch_size,
        shuffle=False,
        num_workers=num_workers,
        pin_memory=True,
    )
    return train_loader, val_loader


# ============================================================
# 3. 训练与评估
# ============================================================

def train_one_epoch(
    model: nn.Module,
    loader: data.DataLoader,
    optimizer: torch.optim.Optimizer,
    device: torch.device,
    epoch: int,
    max_steps: int = 0,
) -> None:
    model.train()
    total_loss = 0.0
    total_correct = 0
    total_seen = 0

    for step, (images, targets) in enumerate(loader, start=1):
        if max_steps > 0 and step > max_steps:
            break

        images = images.to(device, non_blocking=True)
        targets = targets.to(device, non_blocking=True)

        logits = model(images)
        loss = F.cross_entropy(logits, targets)

        optimizer.zero_grad()
        loss.backward()
        optimizer.step()

        with torch.no_grad():
            pred = logits.argmax(dim=1)
            total_correct += (pred == targets).sum().item()
            total_loss += loss.item() * images.size(0)
            total_seen += images.size(0)

    print(
        f"epoch {epoch:03d} | "
        f"loss={total_loss / max(total_seen, 1):.4f} | "
        f"acc={100.0 * total_correct / max(total_seen, 1):.2f}%"
    )


@torch.no_grad()
def evaluate(
    model: nn.Module,
    loader: data.DataLoader,
    device: torch.device,
    title: str = "eval",
    max_batches: int = 0,
) -> Tuple[float, float]:
    model.eval()
    total_loss = 0.0
    total_correct = 0
    total_seen = 0

    for batch_idx, (images, targets) in enumerate(loader, start=1):
        if max_batches > 0 and batch_idx > max_batches:
            break

        images = images.to(device, non_blocking=True)
        targets = targets.to(device, non_blocking=True)

        logits = model(images)
        loss = F.cross_entropy(logits, targets)
        pred = logits.argmax(dim=1)

        total_correct += (pred == targets).sum().item()
        total_loss += loss.item() * images.size(0)
        total_seen += images.size(0)

    avg_loss = total_loss / max(total_seen, 1)
    acc = 100.0 * total_correct / max(total_seen, 1)
    print(f"[{title}] loss={avg_loss:.4f}, acc={acc:.2f}%")
    return avg_loss, acc


# ============================================================
# 4. LeGR: 计算 filter norm 和全局 ranking
# ============================================================

@dataclass
class FilterScore:
    layer_idx: int
    filter_idx: int
    norm2: float
    alpha: float
    kappa: float
    importance: float


def compute_filter_scores(
    model: Tiny5ConvNet,
    alpha: Sequence[float],
    kappa: Sequence[float],
) -> List[FilterScore]:
    """
    计算所有卷积 filters 的 LeGR importance。

    LeGR 公式:
        I_i = alpha_l * ||Theta_i||_2^2 + kappa_l

    其中 l 是 filter 所在层。
    """
    if len(alpha) != 5 or len(kappa) != 5:
        raise ValueError("alpha and kappa must have length 5 for Tiny5ConvNet.")

    scores: List[FilterScore] = []

    for layer_idx, conv in enumerate(model.convs):
        w = conv.weight.detach().cpu()
        # shape: [out_channels, in_channels, k, k]
        norm2 = w.flatten(1).pow(2).sum(dim=1)
        for filter_idx, n2 in enumerate(norm2.tolist()):
            importance = float(alpha[layer_idx]) * float(n2) + float(kappa[layer_idx])
            scores.append(
                FilterScore(
                    layer_idx=layer_idx,
                    filter_idx=filter_idx,
                    norm2=float(n2),
                    alpha=float(alpha[layer_idx]),
                    kappa=float(kappa[layer_idx]),
                    importance=float(importance),
                )
            )

    return scores


def print_global_ranking(scores: List[FilterScore], topk: int = 20) -> None:
    """
    打印全局低重要性 filters。
    """
    ordered = sorted(scores, key=lambda s: s.importance)
    print("\n[Global filter ranking: lowest importance first]")
    print("layer | filter | norm2      | alpha | kappa | importance")
    print("-" * 66)
    for s in ordered[: min(topk, len(ordered))]:
        print(
            f"conv{s.layer_idx + 1:<2d} | "
            f"{s.filter_idx:<6d} | "
            f"{s.norm2:<10.6f} | "
            f"{s.alpha:<5.2f} | "
            f"{s.kappa:<5.2f} | "
            f"{s.importance:<10.6f}"
        )


def summarize_layer_scores(scores: List[FilterScore]) -> None:
    print("\n[Layer-wise score summary]")
    print("layer | channels | norm2_mean | importance_min | importance_mean | importance_max")
    print("-" * 86)
    for layer_idx in range(5):
        layer_scores = [s for s in scores if s.layer_idx == layer_idx]
        norm_vals = torch.tensor([s.norm2 for s in layer_scores])
        imp_vals = torch.tensor([s.importance for s in layer_scores])
        print(
            f"conv{layer_idx + 1:<2d} | "
            f"{len(layer_scores):<8d} | "
            f"{norm_vals.mean().item():<10.6f} | "
            f"{imp_vals.min().item():<14.6f} | "
            f"{imp_vals.mean().item():<15.6f} | "
            f"{imp_vals.max().item():<14.6f}"
        )


# ============================================================
# 5. 根据 global ranking 生成剪枝方案
# ============================================================

@dataclass
class PrunePlan:
    global_prune_ratio: float
    keep_indices: List[List[int]]
    prune_indices: List[List[int]]
    pruned_count: int
    total_filters: int


def make_global_prune_plan(
    model: Tiny5ConvNet,
    scores: List[FilterScore],
    global_prune_ratio: float,
    min_keep: int = 2,
    prune_last: bool = True,
) -> PrunePlan:
    """
    从全局最低 importance filters 开始删除,直到达到 global_prune_ratio。

    这个函数体现 LeGR 的关键思想:
        同一个 global ranking 可以为不同预算生成不同剪枝方案。
    """
    channels = list(model.channels)
    total_filters = sum(channels)
    target_pruned = int(round(total_filters * global_prune_ratio))

    keep_sets = [set(range(c)) for c in channels]
    prune_sets = [set() for _ in channels]

    ordered = sorted(scores, key=lambda s: s.importance)

    for s in ordered:
        if sum(len(x) for x in prune_sets) >= target_pruned:
            break
        if not prune_last and s.layer_idx == 4:
            continue
        if len(keep_sets[s.layer_idx]) <= min_keep:
            continue
        if s.filter_idx not in keep_sets[s.layer_idx]:
            continue

        keep_sets[s.layer_idx].remove(s.filter_idx)
        prune_sets[s.layer_idx].add(s.filter_idx)

    keep_indices = [sorted(list(s)) for s in keep_sets]
    prune_indices = [sorted(list(s)) for s in prune_sets]
    pruned_count = sum(len(x) for x in prune_indices)

    return PrunePlan(
        global_prune_ratio=global_prune_ratio,
        keep_indices=keep_indices,
        prune_indices=prune_indices,
        pruned_count=pruned_count,
        total_filters=total_filters,
    )


def print_prune_plan(plan: PrunePlan) -> None:
    print("\n[LeGR prune plan]")
    print(f"Global prune ratio target: {plan.global_prune_ratio:.2f}")
    print(f"Actually pruned filters: {plan.pruned_count}/{plan.total_filters} "
          f"({100.0 * plan.pruned_count / max(plan.total_filters, 1):.2f}%)")
    for i, (keep, prune) in enumerate(zip(plan.keep_indices, plan.prune_indices), start=1):
        print(
            f"conv{i}: keep {len(keep):2d}, prune {len(prune):2d}, "
            f"pruned indices={prune}"
        )


# ============================================================
# 6. 物理构建剪枝后的窄网络
# ============================================================

@torch.no_grad()
def copy_bn_subset(old_bn: nn.BatchNorm2d, new_bn: nn.BatchNorm2d, keep_out: List[int]) -> None:
    idx = torch.tensor(keep_out, dtype=torch.long)
    new_bn.weight.copy_(old_bn.weight.detach().cpu()[idx])
    new_bn.bias.copy_(old_bn.bias.detach().cpu()[idx])
    new_bn.running_mean.copy_(old_bn.running_mean.detach().cpu()[idx])
    new_bn.running_var.copy_(old_bn.running_var.detach().cpu()[idx])
    new_bn.num_batches_tracked.copy_(old_bn.num_batches_tracked.detach().cpu())


@torch.no_grad()
def build_pruned_model(original: Tiny5ConvNet, plan: PrunePlan) -> Tiny5ConvNet:
    """
    根据 keep_indices 真实构建窄网络。

    对于串行 CNN:
        conv_l 的输出通道被剪
        =
        conv_{l+1} 的输入通道同步被剪

    最后一层 conv5 输出通道被剪后,fc 输入维度同步裁剪。
    """
    new_channels = [len(k) for k in plan.keep_indices]
    pruned = Tiny5ConvNet(channels=new_channels, num_classes=original.num_classes)

    for layer_idx in range(5):
        keep_out = plan.keep_indices[layer_idx]
        out_idx = torch.tensor(keep_out, dtype=torch.long)

        if layer_idx == 0:
            in_idx = torch.arange(3, dtype=torch.long)
        else:
            in_idx = torch.tensor(plan.keep_indices[layer_idx - 1], dtype=torch.long)

        old_conv = original.convs[layer_idx]
        new_conv = pruned.convs[layer_idx]

        # old weight: [old_out, old_in, k, k]
        # new weight: [new_out, new_in, k, k]
        new_weight = old_conv.weight.detach().cpu()[out_idx][:, in_idx, :, :]
        new_conv.weight.copy_(new_weight)

        copy_bn_subset(original.bns[layer_idx], pruned.bns[layer_idx], keep_out)

    final_idx = torch.tensor(plan.keep_indices[-1], dtype=torch.long)
    pruned.fc.weight.copy_(original.fc.weight.detach().cpu()[:, final_idx])
    pruned.fc.bias.copy_(original.fc.bias.detach().cpu())

    return pruned


# ============================================================
# 7. 参数量和 MACs 统计
# ============================================================

def count_params(model: nn.Module) -> int:
    return sum(p.numel() for p in model.parameters())


@torch.no_grad()
def profile_macs(model: nn.Module, device: torch.device, input_size: Tuple[int, int, int, int] = (1, 3, 32, 32)) -> Dict[str, float]:
    model = model.to(device).eval()
    total_macs = 0
    conv_macs = 0
    linear_macs = 0
    handles = []

    def conv_hook(m: nn.Conv2d, _inputs, output):
        nonlocal total_macs, conv_macs
        b, cout, h, w = output.shape
        kernel_ops = m.kernel_size[0] * m.kernel_size[1] * m.in_channels // m.groups
        macs = b * cout * h * w * kernel_ops
        conv_macs += macs
        total_macs += macs

    def linear_hook(m: nn.Linear, inputs, _output):
        nonlocal total_macs, linear_macs
        x = inputs[0]
        b = x.shape[0] if x.dim() > 1 else 1
        macs = b * m.in_features * m.out_features
        linear_macs += macs
        total_macs += macs

    for m in model.modules():
        if isinstance(m, nn.Conv2d):
            handles.append(m.register_forward_hook(conv_hook))
        elif isinstance(m, nn.Linear):
            handles.append(m.register_forward_hook(linear_hook))

    dummy = torch.randn(*input_size, device=device)
    _ = model(dummy)

    for h in handles:
        h.remove()

    return {
        "params": float(count_params(model)),
        "macs": float(total_macs),
        "conv_macs": float(conv_macs),
        "linear_macs": float(linear_macs),
    }


def print_profile(original: nn.Module, pruned: nn.Module, device: torch.device) -> None:
    s0 = profile_macs(original, device)
    s1 = profile_macs(pruned, device)

    def reduction(a: float, b: float) -> float:
        return 0.0 if a == 0 else 100.0 * (a - b) / a

    print("\n[Profile comparison]")
    print(f"Params: {s0['params']:.0f} -> {s1['params']:.0f}, reduction={reduction(s0['params'], s1['params']):.2f}%")
    print(f"MACs:   {s0['macs']:.0f} -> {s1['macs']:.0f}, reduction={reduction(s0['macs'], s1['macs']):.2f}%")
    print(f"Conv:   {s0['conv_macs']:.0f} -> {s1['conv_macs']:.0f}, reduction={reduction(s0['conv_macs'], s1['conv_macs']):.2f}%")
    print(f"Linear: {s0['linear_macs']:.0f} -> {s1['linear_macs']:.0f}, reduction={reduction(s0['linear_macs'], s1['linear_macs']):.2f}%")


# ============================================================
# 8. 简化版 alpha/kappa 搜索
# ============================================================

@dataclass
class AlphaKappaCandidate:
    alpha: List[float]
    kappa: List[float]
    score: float


def random_alpha_kappa(num_layers: int = 5, alpha_std: float = 0.35, kappa_std: float = 0.2) -> Tuple[List[float], List[float]]:
    alpha = [max(0.05, random.gauss(1.0, alpha_std)) for _ in range(num_layers)]
    kappa = [random.gauss(0.0, kappa_std) for _ in range(num_layers)]
    return alpha, kappa


def mutate_alpha_kappa(
    alpha: Sequence[float],
    kappa: Sequence[float],
    mutation_prob: float = 0.5,
    alpha_noise: float = 0.15,
    kappa_noise: float = 0.10,
) -> Tuple[List[float], List[float]]:
    new_alpha = list(alpha)
    new_kappa = list(kappa)
    for i in range(len(new_alpha)):
        if random.random() < mutation_prob:
            new_alpha[i] = max(0.05, new_alpha[i] + random.gauss(0.0, alpha_noise))
        if random.random() < mutation_prob:
            new_kappa[i] = new_kappa[i] + random.gauss(0.0, kappa_noise)
    return new_alpha, new_kappa


def short_finetune(
    model: nn.Module,
    loader: data.DataLoader,
    device: torch.device,
    steps: int,
    lr: float,
) -> None:
    if steps <= 0:
        return
    model.train()
    opt = torch.optim.SGD(model.parameters(), lr=lr, momentum=0.9, weight_decay=5e-4)
    step_count = 0
    while step_count < steps:
        for images, targets in loader:
            if step_count >= steps:
                break
            images = images.to(device, non_blocking=True)
            targets = targets.to(device, non_blocking=True)
            logits = model(images)
            loss = F.cross_entropy(logits, targets)
            opt.zero_grad()
            loss.backward()
            opt.step()
            step_count += 1


def evaluate_alpha_kappa_candidate(
    baseline: Tiny5ConvNet,
    train_loader: data.DataLoader,
    val_loader: data.DataLoader,
    device: torch.device,
    alpha: Sequence[float],
    kappa: Sequence[float],
    target_ratio: float,
    min_keep: int,
    short_steps: int,
    eval_batches: int,
) -> float:
    """
    简化版 LeGR fitness:
        alpha/kappa -> LeGR pruning -> short fine-tune -> validation accuracy
    """
    scores = compute_filter_scores(baseline, alpha, kappa)
    plan = make_global_prune_plan(
        model=baseline,
        scores=scores,
        global_prune_ratio=target_ratio,
        min_keep=min_keep,
        prune_last=True,
    )
    pruned = build_pruned_model(baseline.cpu(), plan).to(device)
    baseline.to(device)

    short_finetune(pruned, train_loader, device, steps=short_steps, lr=0.01)
    _loss, acc = evaluate(pruned, val_loader, device, title="candidate", max_batches=eval_batches)
    return acc


def search_alpha_kappa(
    baseline: Tiny5ConvNet,
    train_loader: data.DataLoader,
    val_loader: data.DataLoader,
    device: torch.device,
    iters: int,
    target_ratio: float,
    min_keep: int,
    short_steps: int,
    eval_batches: int,
    population_size: int = 8,
) -> Tuple[List[float], List[float]]:
    """
    一个很小的 evolutionary search,用于理解 LeGR 如何学习 alpha/kappa。

    注意:这不是论文级 regularized EA 复现,只是教学版本。
    """
    print("\n[Searching alpha/kappa: simplified evolutionary search]")
    print(f"target global prune ratio: {target_ratio:.2f}")

    population: List[AlphaKappaCandidate] = []

    # 初始化候选池
    for _ in range(population_size):
        alpha, kappa = random_alpha_kappa()
        score = evaluate_alpha_kappa_candidate(
            baseline=baseline,
            train_loader=train_loader,
            val_loader=val_loader,
            device=device,
            alpha=alpha,
            kappa=kappa,
            target_ratio=target_ratio,
            min_keep=min_keep,
            short_steps=short_steps,
            eval_batches=eval_batches,
        )
        population.append(AlphaKappaCandidate(alpha, kappa, score))

    population.sort(key=lambda c: c.score, reverse=True)
    print(f"initial best score: {population[0].score:.2f}%")

    for it in range(1, iters + 1):
        # 从前半部分里选 parent,模拟 tournament。
        parent = random.choice(population[: max(1, population_size // 2)])
        alpha, kappa = mutate_alpha_kappa(parent.alpha, parent.kappa)

        score = evaluate_alpha_kappa_candidate(
            baseline=baseline,
            train_loader=train_loader,
            val_loader=val_loader,
            device=device,
            alpha=alpha,
            kappa=kappa,
            target_ratio=target_ratio,
            min_keep=min_keep,
            short_steps=short_steps,
            eval_batches=eval_batches,
        )

        population.append(AlphaKappaCandidate(alpha, kappa, score))
        population.sort(key=lambda c: c.score, reverse=True)
        population = population[:population_size]

        best = population[0]
        print(f"iter {it:03d}/{iters} | new score={score:.2f}% | best={best.score:.2f}%")

    best = population[0]
    print("\nBest alpha:", [round(x, 4) for x in best.alpha])
    print("Best kappa:", [round(x, 4) for x in best.kappa])
    print(f"Best validation score: {best.score:.2f}%")
    return best.alpha, best.kappa


# ============================================================
# 9. 主函数
# ============================================================

def parse_args() -> argparse.Namespace:
    parser = argparse.ArgumentParser(description="LeGR 5-layer CNN global ranking demo")

    parser.add_argument("--data", type=str, default="random", choices=["random", "cifar10"])
    parser.add_argument("--data-root", type=str, default="./data")
    parser.add_argument("--num-train", type=int, default=512)
    parser.add_argument("--num-val", type=int, default=256)
    parser.add_argument("--batch-size", type=int, default=64)
    parser.add_argument("--num-workers", type=int, default=2)

    parser.add_argument("--epochs", type=int, default=1)
    parser.add_argument("--lr", type=float, default=0.05)
    parser.add_argument("--train-max-steps", type=int, default=0)

    parser.add_argument("--alpha", type=float, nargs="+", default=None, help="5 alpha values, one per conv layer.")
    parser.add_argument("--kappa", type=float, nargs="+", default=None, help="5 kappa values, one per conv layer.")
    parser.add_argument("--random-alpha-kappa", action="store_true")

    parser.add_argument("--global-prune-ratios", type=float, nargs="+", default=[0.1, 0.3, 0.5])
    parser.add_argument("--min-keep", type=int, default=2)
    parser.add_argument("--topk", type=int, default=20)

    parser.add_argument("--search-alpha-kappa", action="store_true")
    parser.add_argument("--search-iters", type=int, default=10)
    parser.add_argument("--search-target-ratio", type=float, default=0.4)
    parser.add_argument("--search-short-steps", type=int, default=5)
    parser.add_argument("--search-eval-batches", type=int, default=3)
    parser.add_argument("--population-size", type=int, default=6)

    parser.add_argument("--eval-pruned", action="store_true", help="Evaluate each pruned model.")
    parser.add_argument("--seed", type=int, default=42)
    parser.add_argument("--device", type=str, default="cuda" if torch.cuda.is_available() else "cpu")

    return parser.parse_args()


def resolve_alpha_kappa(args: argparse.Namespace) -> Tuple[List[float], List[float]]:
    if args.alpha is not None:
        if len(args.alpha) != 5:
            raise ValueError("--alpha must contain exactly 5 numbers.")
        alpha = list(map(float, args.alpha))
    else:
        alpha = [1.0] * 5

    if args.kappa is not None:
        if len(args.kappa) != 5:
            raise ValueError("--kappa must contain exactly 5 numbers.")
        kappa = list(map(float, args.kappa))
    else:
        kappa = [0.0] * 5

    if args.random_alpha_kappa:
        alpha, kappa = random_alpha_kappa()

    return alpha, kappa


def main() -> None:
    args = parse_args()
    set_seed(args.seed)
    device = torch.device(args.device)

    print(f"Using device: {device}")
    print(f"Data mode: {args.data}")

    if args.data == "random":
        train_loader, val_loader = build_random_loaders(
            num_train=args.num_train,
            num_val=args.num_val,
            batch_size=args.batch_size,
        )
    else:
        train_loader, val_loader = build_cifar10_loaders(
            data_root=args.data_root,
            batch_size=args.batch_size,
            num_workers=args.num_workers,
            train_subset=args.num_train,
            val_subset=args.num_val,
        )

    model = Tiny5ConvNet().to(device)

    print("\n[Baseline model]")
    print(f"channels: {model.channels}")
    print(f"params: {count_params(model):,}")

    if args.epochs > 0:
        opt = torch.optim.SGD(model.parameters(), lr=args.lr, momentum=0.9, weight_decay=5e-4)
        for epoch in range(1, args.epochs + 1):
            train_one_epoch(model, train_loader, opt, device, epoch, max_steps=args.train_max_steps)

    evaluate(model, val_loader, device, title="baseline", max_batches=0)

    alpha, kappa = resolve_alpha_kappa(args)

    if args.search_alpha_kappa:
        # 搜索时用当前 baseline 模型作为起点。
        alpha, kappa = search_alpha_kappa(
            baseline=model,
            train_loader=train_loader,
            val_loader=val_loader,
            device=device,
            iters=args.search_iters,
            target_ratio=args.search_target_ratio,
            min_keep=args.min_keep,
            short_steps=args.search_short_steps,
            eval_batches=args.search_eval_batches,
            population_size=args.population_size,
        )

    print("\n[LeGR alpha/kappa]")
    print("alpha:", [round(x, 4) for x in alpha])
    print("kappa:", [round(x, 4) for x in kappa])

    scores = compute_filter_scores(model, alpha, kappa)
    summarize_layer_scores(scores)
    print_global_ranking(scores, topk=args.topk)

    original_cpu = copy.deepcopy(model).cpu()

    for ratio in args.global_prune_ratios:
        plan = make_global_prune_plan(
            model=model,
            scores=scores,
            global_prune_ratio=ratio,
            min_keep=args.min_keep,
            prune_last=True,
        )
        print_prune_plan(plan)

        pruned = build_pruned_model(original_cpu, plan)
        print(f"Pruned model channels: {pruned.channels}")
        print_profile(original_cpu, pruned, device)

        if args.eval_pruned:
            evaluate(pruned.to(device), val_loader, device, title=f"pruned ratio={ratio:.2f}")

    print("\nDone.")
    print("Key idea:")
    print("  LeGR does not directly compare raw L2 norms across layers.")
    print("  It learns or sets a layer-wise affine transform alpha/kappa.")
    print("  After transformation, all filters can be globally ranked.")
    print("  One global ranking can generate multiple pruning budgets.")


if __name__ == "__main__":
    main()

五、关键公式

5.1 Filter 全局重要性

I_i = \alpha_{l(i)} \left| \Theta_i \right|2^2 + \kappa{l(i)}

这个公式表示:先计算 filter 的 L2 norm,再根据该 filter 所在层进行缩放和平移,得到可跨层比较的重要性分数。


5.2 Global Ranking 剪枝规则

\text{Prune} = \operatorname{BottomRank} \left( I_1,I_2,\dots,I_K \right)

也就是删除全局重要性最低的 filters,直到满足给定 FLOPs 目标。


5.3 学习 \alpha,\kappa的目标

\arg\max_{\alpha,\kappa} \operatorname{Acc}{val} \left( \hat{\Theta}{l} \right)

其中:

\hat{\Theta}_{l} = \operatorname{LeGR\text{-}Pruning} (\alpha,\kappa,\hat{\zeta}_l)

 

这个目标表示:寻找一组 layer-wise affine transformation 参数,使得在最低 FLOPs 约束下生成的剪枝模型,短暂 fine-tuning 后验证集精度最高。


5.4 Subset Assumption

设  F(f)_l表示 FLOPs 为 (f) 的最优剪枝网络中第 (l) 层保留的 filters 数量,则:

F(f)_l \le F(f')_l, \quad \forall l,\quad \text{if } f \le f'

含义是:更小 FLOPs 模型在每一层保留的 filters 数量都不超过更大 FLOPs 模型。


六、实验设置

LeGR 在多个图像分类 benchmark 上验证,包括:

CIFAR-10
CIFAR-100
ImageNet
Birds-200

论文使用的网络包括:

ResNet-56
VGG-13
ResNet-50
MobileNetV2

CVPR 论文中说明,CIFAR-10/100 各包含 50k 训练图像和 10k 测试图像;ImageNet 包含约 1.2M 训练图像和 50k 测试图像;Birds-200 包含约 6k 训练图像和 5.7k 测试图像。论文也说明,Bird-200 用于 transfer learning 场景分析。

训练设置方面,论文在 CIFAR-10/100 上使用 SGD with Nesterov,weight decay 为5\times 10^{-4} ,batch size 为 128,初始学习率为10^{-1} ,共训练 200 epochs;剪枝后 fine-tuning 在 CIFAR-100 和 Bird-200 上使用较小初始学习率 10^{-2},训练 60 epochs。ImageNet 剪枝模型使用预训练模型并 fine-tune 60 epochs。

LeGR 的搜索设置中,论文使用\hat{\tau}=200  个 gradient updates 作为候选结构搜索阶段的短 fine-tuning,并将搜索架构数设为 400,用于和 AMC 公平比较。


七、实验结果解读

7.1 LeGR 一次搜索生成多种 FLOPs 模型

LeGR 最重要的实验结论不是单个模型精度,而是:学习一次 \alpha,\kappa,就能生成多个 FLOPs 档位的剪枝模型。

例如在 ResNet-56 / CIFAR-100 的实验中,论文使用最低 FLOPs 约束学习\alpha,\kappa ,然后用同一组 \alpha,\kappa生成 20% 到 80% FLOP count 的七个网络。

这正是 LeGR 和 AMC、MorphNet 等方法的核心区别:

AMC / MorphNet:
    每个目标 FLOPs 都要重新搜索一次。

LeGR:
    搜索一次 ranking,
    重复用于多个 FLOPs 目标。

7.2 ResNet-56 / CIFAR-100:低 FLOPs 区间优势明显

论文在 CIFAR-100 上比较了 ResNet-56 和 MobileNetV2。结果显示,LeGR 在低 FLOPs regime 下尤其强。论文指出,在更激进压缩场景中,AMC 和 MorphNet 方差更大,而 LeGR 表现更稳定,并且优于其他方法。

这说明全局 ranking 的优势主要体现在:

不是每层均匀剪,
也不是每个 FLOPs 目标单独搜索,
而是学到一个更稳定的跨层删除顺序。

7.3 剪枝效率:比 AMC / MorphNet 更快

LeGR 的主要卖点之一是效率。

论文比较了获得七个不同 FLOP count 的 ResNet-56 / CIFAR-100 剪枝模型所需时间,并将成本分成:

pruning search cost
fine-tuning cost

结果显示,在 pruning time 上,LeGR 比 AMC 快约 7 倍,比 MorphNet 快约 5 倍。原因是 LeGR 只搜索一次 (\alpha-\kappa) pair,并将它复用于多个 FLOPs 档位;而 AMC 和 MorphNet 需要针对每个 FLOPs 目标重新搜索。

论文摘要也总结 LeGR 在目标为七个不同 accuracy/FLOPs profiles 的 ResNet-56 / CIFAR-100 实验中,比 prior work 快 2× 到 3×,同时性能相当或更好。


7.4 CIFAR-10:ResNet-56 和 VGG-13 上表现有竞争力

在 CIFAR-10 结果中,论文报告 LeGR 在 ResNet-56 上可以从 93.9% fine-tune 到 94.1±0.0%,MFLOPs 为 87.8;在另一个压缩设置下,LeGR 得到 93.7±0.2%,MFLOPs 为 58.9。对于 VGG-13,LeGR 从 91.9% 提升到 92.4±0.2%,MFLOPs 为 70.3。

这个结果说明 LeGR 不仅是效率更高,在单个目标复杂度上的精度也能与已有剪枝方法相当甚至更好。


7.5 ImageNet:ResNet-50 和 MobileNetV2 上验证可扩展性

论文在 ImageNet 上使用 ResNet-50 和 MobileNetV2 验证 LeGR。CVPR 摘要指出,LeGR 在 ImageNet 和 Bird-200 上评估了 ResNet-50 和 MobileNetV2,证明方法的有效性。

更重要的是,论文强调:对于 prior methods,要获得多个不同 FLOPs 的 pruned ConvNets,通常需要为每个 FLOPs 目标分别运行剪枝算法;LeGR 只学习一次 ranking,就能获得多个 FLOPs 档位的模型。


7.6 Bird-200:迁移学习场景下有效

论文还在 Bird-200 上测试 transfer learning 场景。具体做法是先在目标数据集上 fine-tune ImageNet 预训练模型,然后对 fine-tuned network 进行剪枝。论文报告,MobileNetV2 和 ResNet-50 在 Bird-200 上 fine-tune 后分别达到 80.2% 和 79.5% Top-1 accuracy,然后 LeGR 在剪枝结果上优于 Uniform 和 AMC。

这说明 LeGR 不只是标准分类 benchmark 上有效,也适合实际迁移学习中“已有一个目标任务 fine-tuned 模型,希望压缩部署”的场景。


7.7 FLOPs 与真实 runtime

LeGR 的目标是生成不同 speed/accuracy trade-off 的模型,因此论文也讨论了 FLOPs 与 wall-clock runtime 的关系。论文在 ResNet-50 和 MobileNetV2 上,用 PyTorch 0.4 在 Intel i7 和 ARM A57 两类 CPU 上测试,说明 FLOP count 对 runtime 有一定预测作用。

这点对部署导向剪枝很重要:LeGR 虽然主要用 FLOPs 作为复杂度约束,但论文并没有完全忽略真实运行时间,而是尝试验证 FLOPs 和 runtime 的相关性。


八、方法优点

8.1 一次学习 ranking,多次生成模型

这是 LeGR 最大优点。

传统剪枝方法通常是:

每个 FLOPs 目标
    ↓
重新运行一次剪枝算法

LeGR 则是:

学习一次全局 ranking
    ↓
剪到 80% FLOPs
    ↓
剪到 60% FLOPs
    ↓
剪到 40% FLOPs
    ↓
生成一组模型

这使它特别适合探索 accuracy-speed trade-off curve。


8.2 解决了跨层 filter ranking 问题

L1 / L2 norm 可以在同层内部排序,但跨层比较不可靠。LeGR 用 layer-wise affine transformation 学习每层 norm 的缩放和平移,使得不同层 filters 可以放到同一个排序空间中比较。

这比简单 global L1 / global L2 更合理。


8.3 搜索参数少

LeGR 并不直接学习每个 filter 的重要性,而是只学习每层两个参数:

alpha_l
kappa_l

这使搜索空间相对紧凑,也让 EA 搜索更可控。


8.4 对多种网络和数据集有效

LeGR 在 CIFAR-10/100、ImageNet、Bird-200 上测试,并覆盖 ResNet-56、VGG-13、ResNet-50、MobileNetV2 等结构。

这说明它不是只适用于某一个小模型。


8.5 适合部署前探索折中点

实际部署时,我们通常不知道最终该选哪个模型:

精度高但慢?
速度快但精度低?
中间折中?

LeGR 直接输出一组候选模型,让工程人员可以根据真实设备、真实延迟和任务效果选择合适点。这比只输出单一模型更符合实际部署流程。


九、方法局限

9.1 依赖 Subset Assumption

LeGR 的全局 ranking 隐含一个强假设:小模型的保留 filters 是大模型保留 filters 的子集。

这个假设带来效率,但也限制了搜索空间。真实最优 Pareto curve 上的模型未必严格嵌套。例如:

60% FLOPs 最优模型
可能在浅层保留更多通道,
但在深层剪得更多。

80% FLOPs 最优模型
可能有完全不同的层间分配。

LeGR 用一个排序序列生成所有模型,可能错过某些非嵌套最优结构。


9.2 同层内部仍依赖 L2 norm

LeGR 学习的是跨层 affine transformation,但同一层内部的 filter 顺序仍然主要由 L2 norm 决定。

这意味着,如果某一层内部 L2 norm 本身不是好的重要性指标,LeGR 也会受影响。它主要解决的是 inter-layer ranking,不是完全重新定义每个 filter 的重要性。


9.3 搜索仍需要短 fine-tuning

虽然 LeGR 比 AMC / MorphNet 更高效,但搜索 (\alpha,\kappa) 时仍然需要对候选剪枝结构进行短 fine-tuning,并用验证集精度作为 fitness。论文实验中使用 200 个 gradient updates 来近似完整 fine-tuning。

因此,它不是完全 zero-cost pruning。


9.4 仍需最终 fine-tuning

LeGR 得到剪枝结构后,仍然需要 fine-tuning 才能获得最终精度。论文 Figure 2 也明确显示:LeGR-Pruning 得到 filter masks 后,pruned network 需要 fine-tuning。

所以 LeGR 节省的是多目标搜索成本,而不是完全取消训练成本。


9.5 对 Transformer / LLM 不直接适用

LeGR 面向 CNN filter pruning。对于 Transformer、ViT、LLM、VLM,剪枝对象可能是:

attention heads
MLP neurons
tokens
layers
KV cache
vision tokens

LeGR 的思想可以迁移为:

学习跨层 / 跨模块的全局重要性变换,
一次排序生成多个预算模型。

但原始 filter norm + layer-wise affine transformation 不能直接照搬。


十、后续影响

LeGR 的影响主要体现在三个方面。

第一,它把剪枝从“单个复杂度目标”推进到“一条 trade-off curve”。这对真实部署非常重要,因为工程中往往需要比较多个速度—精度折中点,而不是只要一个模型。

第二,它提出了 Learned Global Ranking 这个概念,强调跨层 filter 排名需要学习,而不能简单使用原始 L1 / L2 norm 直接全局排序。

第三,它和 EagleEye 一起代表了剪枝研究中一个新的方向:

不只是问:哪个 filter 重要?
还要问:如何高效搜索和评估一组候选模型?

从专栏脉络看,LeGR 可以放在这里:

Pruning Filters for Efficient ConvNets
    ↓
ThiNet
    ↓
Channel Pruning
    ↓
Network Slimming
    ↓
NISP
    ↓
DCP
    ↓
SFP
    ↓
FPGM
    ↓
HRank
    ↓
GAL
    ↓
EagleEye
    ↓
LeGR
    ↓
Rethinking the Value of Network Pruning

如果说 EagleEye 的问题是:

面对大量候选子网,怎样快速评估哪个更有潜力?

那么 LeGR 的问题是:

能不能学习一个跨层全局 filter 排序,
一次性生成一组不同预算下的剪枝模型?

这就是 LeGR 在剪枝论文脉络中的核心位置。


十一、一句话总结

《Towards Efficient Model Compression via Learned Global Ranking》提出 LeGR,通过学习每层 filter norm 的仿射变换,把同层内可比较的 L2 norm 转换成跨层可比较的全局重要性排序;随后只需剪掉排名靠后的 filters,就能一次生成多个不同 FLOPs / accuracy trade-off 的剪枝模型,从而显著降低多目标剪枝搜索成本。

更多推荐