常量折叠与死代码消除在推理图中的边界

封面信息图

在深度学习模型训练完成并转换为推理计算图后,整张网络中存在着大量的静态先验信息:固定的模型权重(Weights)、固定的偏置项(Biases)、固定的超参数(如 LayerNorm 的 $\epsilon = 1e-5$)、以及在导出时遗留的训练专用分支(如 Dropout、梯度检查点标记)。

对于 AI 编译器前端而言,常量折叠(Constant Folding) 与 死代码消除(Dead Code Elimination, DCE) 是最基础、但也是收益最立竿见影的两项全局优化 Pass。

然而,在面对动态形状(Dynamic Shapes)、量化缩放因子传递以及硬件异常边界时,看似简单的常量折叠与 DCE 也会遇到极其微妙的工程陷阱。

深入理解这两项 Pass 的执行边界,是构建稳健编译优化流水线的前提。

+--------------------------------------------------------------------------+
|                       常量折叠与 DCE 优化转换流水线                          |
+--------------------------------------------------------------------------+
| 原始未优化子图:                                                           |
| %w1 = constant([1.0, 2.0, ...]) : tensor<4096xf32>                       |
| %w2 = constant([0.5, 0.5, ...]) : tensor<4096xf32>                       |
| %w_fused = tensor.mul(%w1, %w2)  <--- 所有输入均为常量!                   |
| %mask = math.greater(%w_fused, 0.0)                                      |
| %unused_branch = custom.training_dropout(%input, 0.1) <--- 无任何后继消费  |
+--------------------------------------------------------------------------+
                                    |
                                    v 运行 Constant Folding & DCE Pass
+--------------------------------------------------------------------------+
| 编译优化后子图:                                                           |
| %w_fused_folded = constant([0.5, 1.0, ...]) <--- 在 Host 编译期直接算好结果|
| // %w1, %w2 节点被作为孤立死代码彻底剔除                                   |
| // %unused_branch 被 DCE 逆向标记彻底剪枝                                 |
+--------------------------------------------------------------------------+

1. 常量折叠(Constant Folding)的触发契约与物理收益

常量折叠的定义非常纯粹:如果一个算子节点的所有输入操作数在编译期都是已知的静态常量(Constant Tensor),那么该算子的计算逻辑完全不需要留在推理期由 GPU 执行,编译器直接在 Host 端 CPU 执行该算子,并用计算产出的静态结果直接替换原算子节点。

典型收益场景:

  1. 权重预处理折叠:在模型加载时,常有针对权重的 Transpose、Reshape、Scale / Quantize 操作。如果不做折叠,每次推理都要对静态权重做一次转置;折叠后,权重在离线阶段就被转置好,推理期完全零开销;
  2. 超参数代数化简:例如将 1.0 / sqrt(variance + 1e-5) 在编译期直接化简为一个常数乘子,将 GPU 运行期的昂贵除法与开方指令,优化为单周期的单条乘法指令。

2. 常量折叠的致命边界:跨平台浮点与异常陷阱

在实现常量折叠 Pass 时,编译器开发者必须极其小心以下两大暗坑:

暗坑一:Host 端与 Target 端的浮点精度漂移(Cross-Compilation Flaw)

如果编译器运行在 x86_64 架构的 Linux 服务器上(Host),而目标推理平台是 ARM 芯片或特定 NPU(Target):

  • x86 CPU 在执行浮点运算时,内部的 FPU 可能会使用 80 位扩展精度,或者在执行 FMA 时采用了不同的舍入模式(Rounding Mode);
  • 如果在 Host 端直接用 C++ std::sqrt 折叠浮点常量,得到的结果可能与目标硬件在运行时计算出的数值在最后几位尾数(Mantissa)上存在差异;
  • 这种微小的差异在累加成千上万层后,可能导致敏感大模型的生成结果出现错乱。
    规避准则:常量折叠的底层数值模拟器必须采用严格遵循 IEEE 754 标准的软浮点库(Soft-Float),或者严格锁定与目标硬件一致的舍入模式。

暗坑二:隐藏的除零与 NaN 传播

如果某个未激活的静态分支中包含除以 0 的计算,在 Host 端折叠时可能会直接触发编译器的浮点异常或生成 NaN。编译器必须具备异常安全探测机制,对非法常数折叠进行保守降级或抛出明确的图校验告警。

3. 死代码消除(DCE)的逆向可达性标记算法

死代码消除的目标是剔除所有对最终输出张量没有任何贡献的孤立算子。

在 SSA 计算图中,DCE 采用经典的逆向图着色/工作表算法(Reverse Worklist Algorithm):

use std::collections::HashSet;

pub fn dead_code_elimination(graph: &mut Graph, outputs: &[ValueId]) {
    let mut live_values: HashSet<ValueId> = outputs.iter().copied().collect();
    let mut worklist: Vec<ValueId> = outputs.to_vec();

    // 1. 从最终输出节点逆向递归标记所有活跃值
    while let Some(val_id) = worklist.pop() {
        if let Some(producer_node_id) = graph.get_producer_node(val_id) {
            let node = graph.get_node(producer_node_id);
            for &input_val in &node.inputs {
                if live_values.insert(input_val) {
                    worklist.push(input_val);
                }
            }
        }
    }

    // 2. 遍历全图,物理删除所有未被标记为活跃的孤立算子节点
    graph.nodes.retain(|node| {
        // 如果该节点产生的所有输出都不在 live_values 中,且该节点无外部副作用(Side-Effect Free)
        let has_live_output = node.outputs.iter().any(|out| live_values.contains(out));
        has_live_output || node.has_side_effects()
    });
}

注意这里关键的 node.has_side_effects():如果一个算子具有外部副作用(例如向磁盘打印日志、修改全局状态机、或向网络发送数据包),即使它的输出张量未被后续算子使用,DCE 也绝对不能将其贸然删除!

通过常量折叠与 DCE 的紧密协同,编译器在最前端为后续的高级算子融合与内存规划扫清了一切干扰,呈现出一张最精纯、最紧凑的计算骨架。

更多推荐