把大模型剪枝写成 Ising 问题:用全局优化决定该删哪些层

2026-09-21 30 预计阅读时间: 1 分钟
来源: huggingface.co AI 摘要 Original link

Disclaimer: This article is an AI-assisted summary. Read it together with the original source when precision matters. The summary may omit context, version differences, or edge cases and is not official documentation.

预计阅读时间:10 分钟

传统剪枝常按单个权重、注意力头或网络层的重要性逐一排序,但这种做法容易忽略一个关键事实:两个模块分别看似可以删除,同时删除却可能让模型性能骤降。把结构化剪枝建模为 Ising 优化问题,价值就在于显式描述这种“模块之间的相互作用”,再全局选择一组满足压缩预算的删除方案。

由于给出的来源信息只有标题,下面不会假定某篇具体工作采用了哪些损失函数或求解器,而是给出一种可以直接实践的建模方式。

从独立打分转向组合决策

假设模型中有一组候选模块,例如 Transformer 层、MLP 块或注意力子模块。定义二进制变量:

  • r_i = 1:删除第 i 个模块;
  • r_i = 0:保留第 i 个模块。

最简单的剪枝方法会为每个模块估计一个独立代价 a_i,然后删除代价最小的若干模块:

[ E_{independent}(r)=\sum_i a_i r_i ]

问题在于,模块之间并不独立。相邻 Transformer 层可能承担相似功能,删除其中任意一层影响有限;但两层一起删除时,误差可能迅速放大。也可能出现相反情况:两个模块高度冗余,同时删除反而比独立估计更划算。

因此可以增加二阶相互作用项:

[ E(r)=\sum_i a_i r_i + \sum_{i<j} b_{ij}r_ir_j ]

这里:

  • a_i 表示单独删除模块 i 的性能代价;
  • b_ij > 0 表示同时删除 ij 会产生额外伤害;
  • b_ij < 0 表示两者可能存在可共同消除的冗余;
  • b_ij ≈ 0 表示暂时可以忽略这对模块的联合影响。

若必须删除固定数量的模块,或者达到指定 FLOPs、显存、参数量预算,可以加入惩罚项:

[ E(r)=\sum_i a_i r_i+\sum_{i<j}b_{ij}r_ir_j+\lambda\left(\sum_i c_i r_i-B\right)^2 ]

c_i 是删除模块 i 能节省的成本,B 是目标压缩量,λ 控制违反预算的惩罚强度。这已经是一个典型的 QUBO(二次无约束二进制优化)形式。

通过变量替换 s_i = 2r_i - 1,其中 s_i ∈ {-1, +1},QUBO 可以转换成 Ising Hamiltonian:

[ H(s)=\sum_i h_i s_i+\sum_{i<j}J_{ij}s_is_j+constant ]

于是,“找出应该删除的模块集合”就变成了“寻找能量最低的自旋配置”。求解并不要求量子计算机:小问题可以穷举,中等问题可以使用模拟退火、局部搜索或整数规划,大问题则需要稀疏化交互矩阵并结合启发式算法。

如何从校准数据估计交互项

一种可以这样实践的估计方法,是在小型校准集上测量负对数似然或困惑度变化。设完整模型的损失为 L0

[ a_i=L(\text{remove }i)-L_0 ]

再测量成对删除后的损失,并扣除两个模块各自的独立影响:

[ b_{ij}=L(\text{remove }i,j)-L_0-a_i-a_j ]

这个差分很重要。若仅把成对删除后的总损失直接当作 b_ij,会重复计算一阶影响。

真实大模型不适合测量所有 O(n²) 个模块对。工程上可以缩小范围:

  1. 只测量相邻层,或距离不超过若干层的组合;
  2. 用激活相似度筛选可能存在冗余的模块对;
  3. 先按一阶代价筛出候选集合,再估计候选内部的交互;
  4. 将绝对值很小的 b_ij 截断为零,得到稀疏图;
  5. 分别在多批校准样本上测量,避免某一领域的数据主导剪枝结果。

测量时还要明确“删除”的实现。残差结构中的一个块可以被替换为恒等映射,但维度变换层、跨层共享参数或带缓存约束的注意力模块未必能直接移除。性能目标也不应只看参数量:实际收益可能来自推理延迟、KV Cache、峰值显存或设备利用率。

一个可运行的最小求解器

下面的 Python 脚本只使用标准库。它假设有 6 个候选块,必须删除其中 2 个,并通过一阶代价与成对交互计算总能量。实际使用时,把 unary_costpair_cost 替换为校准实验测得的数据即可。

from itertools import combinations

# 单独删除每个块时,相对基线增加的验证损失。
unary_cost = {
    0: 0.12,
    1: 0.08,
    2: 0.10,
    3: 0.05,
    4: 0.07,
    5: 0.11,
}

# 同时删除两个块时,超出各自独立代价的额外交互。
# 正值表示组合删除更危险,负值表示可能存在共同冗余。
pair_cost = {
    (0, 1): 0.20,
    (1, 2): 0.15,
    (2, 3): -0.06,
    (3, 4): 0.18,
    (4, 5): -0.04,
    (0, 5): 0.03,
}

blocks_to_remove = 2
block_ids = sorted(unary_cost)


def energy(removed):
    removed = set(removed)
    score = sum(unary_cost[i] for i in removed)

    for (i, j), interaction in pair_cost.items():
        if i in removed and j in removed:
            score += interaction

    return score


ranked = sorted(
    (energy(candidate), candidate)
    for candidate in combinations(block_ids, blocks_to_remove)
)

best_score, best_blocks = ranked[0]
print(f"Best blocks to remove: {best_blocks}")
print(f"Estimated loss increase: {best_score:.4f}")

print("\nTop candidates:")
for score, candidate in ranked[:5]:
    print(f"  remove={candidate}, energy={score:.4f}")

运行方式:

python prune_ising_demo.py

这个例子使用固定删除数量,因此没有显式写入预算惩罚项。若不同模块的计算成本不同,可以遍历满足成本区间的组合,或者把下面的项加进 energy()

penalty = penalty_weight * (saved_cost - target_cost) ** 2
score += penalty

对于几十个候选块,穷举会很快失效。此时可以保持同一个目标函数,替换求解器:

  • 用模拟退火快速取得近似解;
  • 用混合整数规划表达精确预算和互斥约束;
  • 用局部搜索在当前删除集合上执行交换操作;
  • 将模型按层段拆分,先局部求解,再做全局修正。

落地时不要把“最低能量”当成最终答案

Ising 或 QUBO 只是搜索代理目标,最终仍需回到真实模型验证。推荐采用以下闭环:

  1. 在校准集上估计单模块和关键模块对的损失变化;
  2. 在硬件指标约束下求出多个低能量候选,而不是只保留一个解;
  3. 真正修改模型结构,并重新测量困惑度、任务准确率和生成质量;
  4. 在目标硬件上测试吞吐量、首 Token 延迟和峰值显存;
  5. 对候选模型进行短程微调或蒸馏,再比较恢复后的质量;
  6. 用新的测量结果更新 a_ib_ij,必要时再次求解。

还需要警惕三个边界:二阶模型无法完整描述三个以上模块的高阶耦合;校准集过小会产生噪声较大的交互项;代理损失最低的结构不一定拥有最低的实际延迟。更稳妥的做法,是把 Ising 优化当作候选结构生成器,而不是替代端到端评测的裁判。

当独立重要性排序开始频繁选出“单看合理、组合糟糕”的剪枝方案时,引入成对交互就很有价值。它把剪枝从一张静态排行榜,升级为一个带预算、依赖关系和硬件目标的组合优化问题。


相关推荐