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

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.05 05:01:19