jax.ops 模块

jax.ops 模块#

函数 jax.ops.index_updatejax.ops.index_add 等,在 JAX 0.2.22 中已弃用,现已移除。请改用 JAX 数组上的 jax.numpy.ndarray.at 属性。

段缩减运算符#

segment_max(data, segment_ids[, ...])

计算数组段内的最大值。

segment_min(data, segment_ids[, ...])

计算数组段内的最小值。

segment_prod(data, segment_ids[, ...])

计算数组段内的乘积。

segment_sum(data, segment_ids[, ...])

计算数组段内的总和。