如何基于二维坐标统计并更新JAX二维数组的计数?
问题描述
现有如下JAX代码:
import jax.numpy as jnp x = jnp.zeros((5,5)) coords = jnp.array([ [1,2], [2,3], [1,2], ])
需要统计coords中每个(行, 列)坐标在x中的出现次数,目标输出为:
Array([[0., 0., 0., 0., 0.], [0., 0., 2., 0., 0.], [0., 0., 0., 1., 0.], [0., 0., 0., 0., 0.], [0., 0., 0., 0., 0.]], dtype=float32)
尝试使用x.at[coords].add(1)得到错误结果:
Array([[0., 0., 0., 0., 0.], [2., 2., 2., 2., 2.], [3., 3., 3., 3., 3.], [1., 1., 1., 1., 1.], [0., 0., 0., 0., 0.]], dtype=float32)
解决方案
方法1:拆分行列索引使用at操作
JAX中,直接用(N,2)的数组作为at的索引会被解析为行索引,而非(行,列)的二维索引。正确做法是将坐标拆分为行、列两个一维数组,分别传入索引:
import jax.numpy as jnp x = jnp.zeros((5,5)) coords = jnp.array([ [1,2], [2,3], [1,2], ]) # 拆分行、列索引 rows = coords[:, 0] cols = coords[:, 1] # 对指定行列位置累加计数 result = x.at[rows, cols].add(1) print(result)
该方法会自动处理重复坐标的累加,直接得到目标结果。
方法2:使用jnp.bincount统计次数
将二维坐标转换为一维索引,通过bincount统计每个索引的出现次数,再重塑为原数组形状:
import jax.numpy as jnp x = jnp.zeros((5,5)) coords = jnp.array([ [1,2], [2,3], [1,2], ]) # 将二维坐标转为一维索引 flat_indices = coords[:, 0] * x.shape[1] + coords[:, 1] # 统计每个索引的出现次数,minlength确保覆盖所有位置 counts = jnp.bincount(flat_indices, minlength=x.size) # 重塑为原数组形状并转换类型 result = counts.reshape(x.shape).astype(x.dtype) print(result)
内容的提问来源于stack exchange,提问作者oneloop
相关产品推荐
相关产品推荐

