JAX中如何更新二维数组的整行或整列?
在JAX中实现整行/整列的加法更新(支持JIT编译)
更新整行
你不需要显式循环,JAX的at索引API支持直接对整行进行加法操作,用add()方法替代set()即可,逻辑和NumPy写法一致:
import jax.numpy as jnp # 示例:给第3行(索引2)添加行向量r mat = jnp.ones((3, 3)) r = jnp.array([1, 2, 3]) updated_mat = mat.at[2].add(r)
mat.at[2]会选中矩阵的第2行,add(r)直接将行向量r与该行元素逐元素相加,返回更新后的新数组(JAX数组不可变,不会原地修改原数组)。
更新整列
更新列的逻辑类似,只需用切片选中目标列([:, j]表示所有行的第j列),再调用add():
# 示例:给第2列(索引1)添加列向量c c = jnp.array([[1], [2], [3]]) # shape (3,1) updated_mat = mat.at[:, 1].add(c.squeeze()) # 转为shape (3,)后相加 # 或者直接利用JAX广播机制,无需squeeze: updated_mat = mat.at[:, 1].add(c)
JAX的广播机制会自动处理(n,1)和(n,)形状的兼容,两种写法都能正常工作。
关键说明
- 这些
at方法的操作都是可JIT编译的,完全符合JAX的纯函数要求,可直接用jax.jit()装饰包含该逻辑的函数。 - JAX数组是不可变对象,所有
at操作都会返回新数组,不会修改原数组,这一点和NumPy的原地修改不同,需要注意。
内容的提问来源于stack exchange,提问作者CartesianBear
相关产品推荐
相关产品推荐

