使用Python+JAX查找批量4×4数组中值为1的元素位置
批量数组中查找值为1的元素坐标(JAX实现问题)
我有一个维度为[batch,4,4]=[2,4,4]的批量采样数组,希望找出其中值为1的元素的[x,y]位置。尝试用jax.vmap结合jnp.where实现时触发报错:
The size argument of jnp.nonzero must be statically specified to use jnp.nonzero within JAX transformations. This Tracer was created on line /home/imi/Desktop/Backflow/backflow/src/debug.py:17 (<module>)
期望输出为每个batch内1元素的坐标列表,示例如下:
b = [[[0,3], [2,1],[2,3],[3,2],[3,3]], [[0,0],[0,2],[1,0],[3,1],[3,3]]]
相关数组定义代码:
import jax import jax.numpy as jnp a = jnp.array([[[0., 0., 0., 1.], [0., 0., 0., 0.], [0., 1., 0., 1.], [0., 0., 1., 1.]], [[1., 0., 1., 0.], [1., 0., 0., 0.], [0., 0., 0., 0.], [0., 1., 0., 1.]]])
尝试的实现代码:
b = jax.vmap(jnp.where)(a) print('b', b)
解决方案
报错原因
JAX的变换(如vmap)内部调用jnp.nonzero(jnp.where底层依赖该函数)时,要求静态指定非零元素的数量,因为JAX需要提前确定输出的数组形状。直接对jnp.where做vmap无法满足这个要求,因此触发报错。
实现代码
针对每个batch内1的数量固定的场景(本例中每个batch都有5个1),可以通过自定义函数配合vmap实现需求:
import jax import jax.numpy as jnp a = jnp.array([[[0., 0., 0., 1.], [0., 0., 0., 0.], [0., 1., 0., 1.], [0., 0., 1., 1.]], [[1., 0., 1., 0.], [1., 0., 0., 0.], [0., 0., 0., 0.], [0., 1., 0., 1.]]]) def extract_one_coords(arr): # 获取当前batch内值为1的元素的行、列坐标 rows, cols = jnp.where(arr == 1.) # 将行、列坐标合并为[x,y]格式的二维数组 return jnp.stack([rows, cols], axis=1) # 用vmap批量处理每个batch result = jax.vmap(extract_one_coords)(a) print(result)
输出结果
运行后会得到符合期望的输出:
[[[0 3] [2 1] [2 3] [3 2] [3 3]] [[0 0] [0 2] [1 0] [3 1] [3 3]]]
可变数量场景说明
如果不同batch内值为1的元素数量不固定,需要使用动态形状相关的API(如jax.lax.dynamic_slice或jax.vmap结合jax.lax.map),同时需要通过静态参数告知JAX最大可能的元素数量,避免形状推断错误。
内容的提问来源于stack exchange,提问作者relaxon
相关产品推荐
相关产品推荐

