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编译器在融合优化与值复用逻辑之间的权衡导致的,具体原因如下:
- vmap展开与融合优化的边界限制
当vmap处理blackbox_function时,编译器会先将函数逻辑展开为对y每个元素的映射计算。为了提升批量计算效率,编译器会把与y元素相关的计算打包到%fused_computation融合块中,但这个融合块的设计目标是输出批量处理后的y结果,内部计算的x+1中间值并没有被暴露为融合块的输出。
后续编译器识别到你只需要返回x[0](即x+1的标量值),但因为融合块被视为一个相对独立的计算单元,跨块的公共子表达式消除(CSE)优化没有触发——编译器默认不会主动提取融合块内部的中间结果供外部复用,除非能明确证明这种复用的收益远大于额外的通信开销。
HLO代码与实际执行的差异
需要注意的是,HLO代码是编译器优化过程中的中间表示,最终生成的硬件指令可能会进一步优化掉重复的加法操作。你可以通过jax.profiler监控实际硬件执行情况,验证是否真的执行了两次加法。手动优化方案
如果想强制实现只计算一次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
相关产品推荐
相关产品推荐

