JAX编译方法中如何实现可变轴的数组滚动操作?
为啥报错?
jnp.roll的axis参数必须是编译时就能确定的静态值,但你用jax.lax.map遍历indList时,每个ind都是动态变化的张量,JAX编译阶段无法确定具体值,因此触发ConcretizationTypeError。给roll加@partial(jax.jit, static_argnums=0)没用,因为map会把ind动态传递给roll,满足不了jit对静态参数的要求——静态参数得在调用时就明确是固定值,不能是动态遍历的变量。
怎么解决?
要实现动态轴的滚动操作,得绕开jnp.roll对静态axis的限制,下面给两种可行的方法:
方法1:手动拼数组实现滚动
不用jnp.roll,直接把数组切分成两部分再拼接,用jax.lax.dynamic_slice_in_dim处理动态轴的切分:
import jax.numpy as jnp from jax import lax A = jnp.ones((4, 4, 4, 4)) indList = jnp.asarray([0, 1, 2]) def dynamic_roll(x, axis, shift): # 把负偏移转成正的,避免越界 shift = shift % x.shape[axis] # 按偏移量切分数组 part1 = lax.dynamic_slice_in_dim(x, shift, x.shape[axis] - shift, axis=axis) part2 = lax.dynamic_slice_in_dim(x, 0, shift, axis=axis) return jnp.concatenate([part1, part2], axis=axis) # 用map遍历每个轴,执行动态滚动 result = lax.map(lambda ind: dynamic_roll(A, ind, -1), indList) print(result.shape) # 输出 (3, 4, 4, 4, 4),符合预期
方法2:枚举轴的情况用switch选择
如果你的轴范围固定(比如例子里只有0、1、2),可以用jax.lax.switch根据动态轴的值,选择对应静态轴的jnp.roll操作:
import jax.numpy as jnp from jax import lax A = jnp.ones((4, 4, 4, 4)) indList = jnp.asarray([0, 1, 2]) def roll_by_axis(ind): # 按ind的值选择对应轴的滚动操作 return lax.switch(ind, [ lambda: jnp.roll(A, -1, axis=0), lambda: jnp.roll(A, -1, axis=1), lambda: jnp.roll(A, -1, axis=2) ]) result = lax.map(roll_by_axis, indList) print(result.shape) # 输出 (3, 4, 4, 4, 4)
内容的提问来源于stack exchange,提问作者rak
相关产品推荐
相关产品推荐

