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

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.09.24 16:36:03