You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

JIT编译的JAX函数存在不必要的重复计算?原因探究

问题:JAX JIT编译器重复执行加法操作的原因分析

最小可复现示例(MWE)

def blackbox_function(x, y): # 给定的黑盒函数,无法拆分
    x = x + 1 # *保证* x的更新方式与y无关
    y = x * y
    return x, y

@jax.jit
def my_mapped_function(x, y):
    x, y = jax.vmap(blackbox_function, in_axes=(None,0))(x, y) # 并行化时x固定,所有y对应的x输出一致
    return x[0], y # 仅返回x的第一个元素

x = 1
y = jnp.arange(10)
print(my_mapped_function.lower(x, y).compile().as_text()) # 查看编译器处理逻辑

流程特性

  • 输入x的更新方式固定,与y无关;
  • 仅返回映射后x的第一个元素(所有元素值一致)。

理论上最优计算应仅执行一次x = x+1并将其作为第一个输出参数,但查看编译后的HLO代码发现加法被执行了两次:

编译后的HLO代码

HloModule jit_my_mapped_function, is_scheduled=true, entry_computation_layout={(s64[], s64[10]{0})->(s64[], s64[10]{0})}, allow_spmd_sharding_propagation_to_parameters={true,true}, allow_spmd_sharding_propagation_to_output={true,true}

%fused_computation (param_0.1: s64[10], param_1.2: s64[]) -> s64[10] {
  %param_1.2 = s64[] parameter(1)
  %constant.0 = s64[] constant(1)
  %add.0 = s64[] add(%param_1.2, %constant.0), metadata={op_name="jit(my_mapped_function)/jit(main)/add" source_file="/var/folders/x0/28x522xx1vb2xl75tn781lqr0000gn/T/ipykernel_33608/2232785945.py" source_line=2}
  %broadcast.3 = s64[10]{0} broadcast(%add.0), dimensions={}, metadata={op_name="jit(my_mapped_function)/jit(main)/mul" source_file="/var/folders/x0/28x522xx1vb2xl75tn781lqr0000gn/T/ipykernel_33608/2232785945.py" source_line=3}
  %param_0.1 = s64[10]{0} parameter(0)
  ROOT %multiply.0 = s64[10]{0} multiply(%broadcast.3, %param_0.1), metadata={op_name="jit(my_mapped_function)/jit(main)/mul" source_file="/var/folders/x0/28x522xx1vb2xl75tn781lqr0000gn/T/ipykernel_33608/2232785945.py" source_line=3}
}

ENTRY %main.11 (Arg_0.1: s64[], Arg_1.2: s64[10]) -> (s64[], s64[10]) {
  %Arg_0.1 = s64[] parameter(0), metadata={op_name="x"}
  %Arg_1.2 = s64[10]{0} parameter(1), metadata={op_name="y"}
  %constant.3 = s64[] constant(1)
  %broadcast_multiply_fusion = s64[10]{0} fusion(%Arg_1.2, %Arg_0.1), kind=kLoop, calls=%fused_computation, metadata={op_name="jit(my_mapped_function)/jit(main)/mul" source_file="/var/folders/x0/28x522xx1vb2xl75tn781lqr0000gn/T/ipykernel_33608/2232785945.py" source_line=3}
  %add.4 = s64[] add(%Arg_0.1, %constant.3), metadata={op_name="jit(my_mapped_function)/jit(main)/add" source_file="/var/folders/x0/28x522xx1vb2xl75tn781lqr0000gn/T/ipykernel_33608/2232785945.py" source_line=2}
  ROOT %tuple.10 = (s64[], s64[10]{0}) tuple(%add.4, %broadcast_multiply_fusion)
}

疑惑点

  • 为何x=x+1的加法操作会被执行两次?
  • 既然融合计算%fused_computation中已计算过一次,为何不直接复用该结果而非要重新计算?
  • 是否存在对HLO代码的误读?

解答

你的解读没有问题,HLO代码里确实存在两次加法操作的表述,这是JAX编译器在融合优化与值复用逻辑之间的权衡导致的,具体原因如下:

  1. vmap展开与融合优化的边界限制
    当vmap处理blackbox_function时,编译器会先将函数逻辑展开为对y每个元素的映射计算。为了提升批量计算效率,编译器会把与y元素相关的计算打包到%fused_computation融合块中,但这个融合块的设计目标是输出批量处理后的y结果,内部计算的x+1中间值并没有被暴露为融合块的输出。

后续编译器识别到你只需要返回x[0](即x+1的标量值),但因为融合块被视为一个相对独立的计算单元,跨块的公共子表达式消除(CSE)优化没有触发——编译器默认不会主动提取融合块内部的中间结果供外部复用,除非能明确证明这种复用的收益远大于额外的通信开销。

  1. HLO代码与实际执行的差异
    需要注意的是,HLO代码是编译器优化过程中的中间表示,最终生成的硬件指令可能会进一步优化掉重复的加法操作。你可以通过jax.profiler监控实际硬件执行情况,验证是否真的执行了两次加法。

  2. 手动优化方案
    如果想强制实现只计算一次x+1,可以重构代码提前计算更新后的x,再传入vmap:

@jax.jit
def my_mapped_function(x, y):
    x_updated = x + 1
    def wrapped_blackbox(y_val):
        return x_updated * y_val
    y = jax.vmap(wrapped_blackbox)(y)
    return x_updated, y

这种写法能让编译器自然只计算一次x+1,同时保证逻辑与原代码完全一致。

内容的提问来源于stack exchange,提问作者Ben

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.06.01 18:39:50