JAX:如何规避单轴元素数量变化时JIT函数的重新编译行为
JAX JIT动态轴避免重编译解决方案
结论
可以避免重新编译,只需将发生变化的轴显式标记为动态维度即可。
重编译原因
默认配置下,jax.jit会把输入张量的所有维度的具体数值作为编译缓存的键值,只要任意维度大小发生变化,就会触发函数重新编译。你示例中输入张量的第0维大小随输入数据变化,因此调用c时触发了重编译。
解决方法
使用jax.jit的in_shapes参数,将需要动态变化的轴标注为None,告知JAX该维度的大小不需要固化到编译逻辑中,修改后的代码如下:
import jax import jax.numpy as jnp # 标注输入参数x的形状为(None, 2, 2):第0维为动态大小,后两维固定为2,dtype匹配输入的整数类型 @jax.jit(in_shapes=jax.ShapeDtypeStruct(shape=(None, 2, 2), dtype=jnp.int32)) def f(x): print('recompiling') return (x + 10) * 100 # 首次调用触发一次编译 a = f(jnp.arange(300000000).reshape((-1, 2, 2)).block_until_ready()) # 相同形状输入,不重编译 b = f(jnp.arange(300000000).reshape((-1, 2, 2)).block_until_ready()) # 第0维变化,其余维度匹配,不重编译 c = f(jnp.arange(450000000).reshape((-1, 2, 2)).block_until_ready())
注意事项
- 被标记为动态的维度,不能作为静态参数在函数中使用,比如不能将动态轴的大小传入需要静态常量参数的JAX接口(如
jax.lax.fori_loop的循环次数参数),如果需要获取动态轴的运行时大小,可以使用d0 = jax.lax.dynamic_shape(x)[0][0]的写法 - 如果函数有多个输入,可以给
in_shapes传入元组,逐个配置每个输入的动态维度规则 - 该功能需要jax版本>=0.4.0,建议升级到最新稳定版使用
内容的提问来源于stack exchange,提问作者mutableVoid
相关产品推荐
相关产品推荐

