JAX 段归约操作指南:jax.ops.segment_sum / segment_max 系列函数与 .at 索引更新全面解析
【免费下载链接】jaxComposable transformations of Python+NumPy programs: differentiate, vectorize, JIT to GPU/TPU, and more项目地址: https://gitcode.com/GitHub_Trending/ja/jax
本文基于 JAX 官方 API 文档 docs/jax.ops.rst 展开,系统讲解jax.ops模块中段归约(Segment Reduction)操作符segment_sum、segment_prod、segment_min、segment_max的完整用法、全部参数语义与底层 scatter 实现原理,同时梳理旧版index_update系列函数的弃用迁移路径(改用jax.numpy.ndarray.at属性)。读完本文,你将能够在数据处理、图聚合、稀疏归约等场景中正确、高效地使用段归约算子,并理解其在 JIT 编译与自动微分下的约束与最佳实践。
一、jax.ops 模块:定位与公开 API
jax.ops是 JAX 中存放“基于索引的操作(indexed operations)”的顶层模块,官方文档通过automodule自动收集其 docstring 与成员签名。从当前仓库的 jax/ops/init.py 可以看到,模块当前对外公开的完整 API 恰为四个段归约函数:
from jax._src.ops.scatter import ( segment_sum as segment_sum, segment_prod as segment_prod, segment_min as segment_min, segment_max as segment_max, )值得注意的是,文件中有一行注释:# Note: import <name> as <name> is required for names to be exported.(PEP 484 命名导出要求),这解释了为什么模块中采用“别名同名导入”的写法。也就是说,当前jax.ops的对外契约就是这四个段归约算子;其余历史上曾经属于jax.ops的函数(如index_update系列)已全部移除,见下文第二节。
二、已移除的 index_update 系列:迁移到 .at 属性
文档明确记载:
The functions
jax.ops.index_update,jax.ops.index_add, etc., which were deprecated in JAX 0.2.22, have been removed. Please use thejax.numpy.ndarray.atproperty on JAX arrays instead.
即:jax.ops.index_update、jax.ops.index_add等函数在 JAX 0.2.22 中标记弃用,并在后续版本中彻底移除。任何新代码都应改用 JAX 数组的.at属性进行“函数式原地更新”。
2.1.at属性的功能对照
在 jax/_src/numpy/array_methods.py 中,_IndexUpdateHelper类的 docstring 给出了一组与 NumPy 原地表达式的等价对照表:
.at语法 | 等价的 NumPy 原地表达式 |
|---|---|
x = x.at[idx].set(y) | x[idx] = y |
x = x.at[idx].add(y) | x[idx] += y |
x = x.at[idx].subtract(y) | x[idx] -= y |
x = x.at[idx].multiply(y) | x[idx] *= y |
x = x.at[idx].divide(y) | x[idx] /= y |
x = x.at[idx].power(y) | x[idx] **= y |
x = x.at[idx].min(y) | x[idx] = minimum(x[idx], y) |
x = x.at[idx].max(y) | x[idx] = maximum(x[idx], y) |
x = x.at[idx].apply(ufunc) | ufunc.at(x, idx) |
x = x.at[idx].get() | x = x[idx] |
该 docstring 特别强调两个与 NumPy 的关键差异:
- 纯函数性:任何
x.at[...]表达式都不会修改原数组x,而是返回修改后的副本;不过在jax.jit编译的函数内部,x = x.at[idx].set(y)这类表达式保证会被原地应用(buffer donation 优化)。 - 重复索引语义:与 NumPy 的
x[idx] += y(只保留最后一次更新)不同,.at会应用所有更新;冲突更新的应用顺序是“实现定义”的,在某些硬件平台上可能因并发而具有不确定性。
2.2.at的越界与模式参数
_IndexUpdateHelper的 docstring 还定义了越界索引的处理模式(mode参数):
"promise_in_bounds"(默认):用户承诺索引在界内,不做额外检查;实际行为是get()中的越界索引被裁剪(clip),set()/add()等更新中的越界索引被丢弃(drop)。"clip":将越界索引裁剪到合法范围。"drop":忽略越界索引。"fill":"drop"的别名;对get()而言,可通过可选参数fill_value指定越界返回值(默认对非精确类型为NaN、有符号类型为最大负值、无符号类型为最大正值、布尔为True)。
此外还有wrap_negative_indices(默认True,负索引从数组末尾计数;设为False时负索引按越界处理)、indices_are_sorted与unique_indices(见 4.3 节)。示例(摘自源码 docstring):
>>> x = jnp.arange(5.0) >>> x.at[2].add(10) Array([ 0., 1., 12., 3., 4.], dtype=float32) >>> x.at[10].add(10) # 默认模式:越界更新被丢弃 Array([0., 1., 2., 3., 4.], dtype=float32) >>> x.at[20].add(10, mode='clip') # clip 模式:裁剪到末尾 Array([ 0., 1., 2., 3., 14.], dtype=float32) >>> x.at[20].get() # get() 默认裁剪 Array(4., dtype=float32) >>> x.at[20].get(mode='fill', fill_value=-1) Array(-1., dtype=float32)在实现层面,.at[idx].set/add/...返回的_IndexUpdateRef对象(array_methods.py#L1150)最终都汇聚到jax._src.ops.scatter模块的_scatter_update帮助函数,再分别映射到lax_slicing.scatter、scatter_add、scatter_mul、scatter_min、scatter_max等底层原语——这正是下一节段归约算子的实现基础。
三、Segment Reduction 段归约算子:核心用法
段归约的目标是:给定一维整数数组segment_ids与数据数组data,将data沿其首轴按照segment_ids划分成若干段,并对每一段施加归约操作(求和、求积、取最大、取最小),输出形状为(num_segments,) + data.shape[1:]。
四个函数签名完全一致(见 jax/_src/ops/scatter.py 第 221、279、339、398 行):
segment_sum(data, segment_ids, num_segments=None, indices_are_sorted=False, unique_indices=False, bucket_size=None, mode=None, out_sharding=None) segment_prod(data, segment_ids, num_segments=None, indices_are_sorted=False, unique_indices=False, bucket_size=None, mode=None, out_sharding=None) segment_max(data, segment_ids, num_segments=None, indices_are_sorted=False, unique_indices=False, bucket_size=None, mode=None, out_sharding=None) segment_min(data, segment_ids, num_segments=None, indices_are_sorted=False, unique_indices=False, bucket_size=None, mode=None, out_sharding=None)3.1 基础示例(摘自源码 docstring)
from jax import jit import jax.numpy as jnp from jax.ops import segment_sum, segment_prod, segment_max, segment_min # --- segment_sum:段求和 --- data = jnp.arange(5) segment_ids = jnp.array([0, 0, 1, 1, 2]) segment_sum(data, segment_ids) # Array([1, 5, 4], dtype=int32) # [0+1, 2+3, 4] # --- segment_prod:段求积 --- data = jnp.arange(6) segment_ids = jnp.array([0, 0, 1, 1, 2, 2]) segment_prod(data, segment_ids) # Array([ 0, 6, 20], dtype=int32) # [0*1, 2*3, 4*5] # --- segment_max:段最大值 --- segment_max(data, segment_ids) # Array([1, 3, 5], dtype=int32) # --- segment_min:段最小值 --- segment_min(data, segment_ids) # Array([0, 2, 4], dtype=int32)3.2 多维数据
segment_ids的长度必须等于data.shape[0];归约只发生在首轴,其余轴原样保留:
data = jnp.arange(12).reshape(6, 2) # data = [[0,1],[2,3],[4,5],[6,7],[8,9],[10,11]] segment_ids = jnp.array([0, 0, 1, 1, 2, 2]) segment_sum(data, segment_ids) # shape (3, 2) # Array([[ 2, 4], [10, 12], [18, 20]], dtype=int32)四、参数语义深度解析
4.1num_segments:输出段数(JIT 下必须为静态值)
- 默认值为
None,此时按max(segment_ids) + 1自动推导(源码 scatter.py#L191-L192)。 - 由于
num_segments直接决定输出张量的形状,在 JIT 编译的函数中使用时必须显式传入静态(concrete)值。源码通过core.concrete_dim_or_error(num_segments, ...)强制校验,非静态值会直接报错。 - 传入负数会抛出
ValueError("num_segments must be non-negative.")。
jit(segment_sum, static_argnums=2)(data, segment_ids, 3) # Array([1, 5, 4], dtype=int32)4.2mode:越界段 id 的处理
默认mode=None,实际映射为GatherScatterMode.FILL_OR_DROP(源码 scatter.py#L187),即:落在[0, num_segments)范围之外的索引被丢弃,不参与归约。该参数接受jax.lax.GatherScatterMode枚举值或其字符串形式("clip"、"fill"、"drop"、"promise_in_bounds"等),语义与 2.2 节.at的 mode 一致。
测试 tests/lax_numpy_indexing_test.py#L1817-L1829 验证了越界与负段 id 的行为(注意负索引默认会被规范化):
data = jnp.array([5, 1, 7, 2, 3, 4, 1, 3]) segment_ids = jnp.array([0, 4, 8, 1, 2, -6, -1, 3]) segment_sum(data, segment_ids, num_segments=4) # Array([5, 2, 3, 3]) # 越界 id(4、8)与规范化后的负 id 按 mod 折回/丢弃4.3indices_are_sorted与unique_indices:性能提示
indices_are_sorted=True:声明segment_ids(规范化后)升序排列。若声明与实际不符,输出未定义。unique_indices=True:声明每个段 id 至多出现一次。若声明与实际不符,输出未定义。
这两项本质上是把“保证”交给用户、换取部分后端更高效的执行路径(源码在 scatter.py#L145-L147 将其并入底层 scatter 调用)。同时注意:segment_prod/segment_max/segment_min的自动微分只在unique_indices=True时完整实现——例如测试 tests/lax_numpy_indexing_test.py#L1782-L1785 断言scatter_mul的梯度在非唯一索引下会抛出NotImplementedError(“scatter_mul gradients are only implemented ifunique_indices=True”)。
4.4bucket_size:数值稳定性分桶
默认None表示不分桶。传入正整数时,算法会把segment_ids按顺序切成多个桶(每桶最多bucket_size个元素),在每个桶内单独执行段归约,再对桶间结果做二次归约(reducer(out, axis=0))。这样做的动机是改善 sum/prod 这类累积归约的数值稳定性(源码注释 scatter.py#L205-L206:"Bucketize indices and perform segment_update on each bucket to improve numerical stability for operations like product and sum")。桶数由num_buckets = ceil(segment_ids.size / bucket_size)决定。
4.5out_sharding:单程序多数据(SPMD)分片
四个函数均支持out_sharding参数(接受NamedSharding或PartitionSpec),用于指定输出的分片布局。源码先通过canonicalize_sharding(out_sharding, 'segment_xxx')规范化,再沿网格显式轴调用auto_axes(scatter.py#L86-L90)。两点实现限制值得注意:
- 传入
out_sharding时不能再同时使用bucket_size(scatter.py#L208-L209 直接raise NotImplementedError); segment_prod、segment_max、segment_min对unreduced(未归约)分片规格尚未支持(scatter.py#L331-L332 等)。
分片用法示例可参考 tests/pjit_test.py#L11201-L11247 中test_segment_sum/test_segment_max/test_segment_prod的写法(在 mesh 上指定out_sharding=out_s并配合jax.jit使用)。
五、底层实现原理:一切归约为 scatter
5.1 共享的_segment_update流水线
四个函数都是薄封装,最终汇聚到同一个私有函数_segment_update(scatter.py#L175-L218),区别仅在于底层 scatter 原语与归约器:
| 公开函数 | 底层 scatter 原语 | 归约器 | 恒等元(identity) |
|---|---|---|---|
segment_sum | scatter_add | sum | 0 |
segment_prod | scatter_mul | prod | 1 |
segment_min | scatter_min | min | +inf(整数为类型最大值,布尔为True) |
segment_max | scatter_max | max | -inf(整数为类型最小值,布尔为False) |
执行流程为:
- 校验输入(
check_arraylike)、将data与segment_ids转为 JAX 数组; - 确定
num_segments(默认max(segment_ids)+1)并做静态性校验; - 用
_get_identity(op, dtype)求得该归约的恒等元(scatter.py#L153-L172:scatter_min对布尔取True、整数取iinfo(dtype).max、浮点取inf,scatter_max反之); - 以恒等元构造形状为
(num_segments,) + data.shape[1:]的输出缓冲区out; - 调用
_scatter_update(out, segment_ids, data, scatter_op, ...),把data按段 id散落到输出缓冲区——这正是 XLA scatter 语义的体现:源码注释明确写道 “XLA gathers and scatters are very similar in structure; the scatter logic is more or less a transpose of the gather equivalent”(scatter.py#L76-L77)。
_scatter_update内部会把用户索引规范化为NDIndexer(支持整数、切片、省略号、布尔数组等高级索引),再转换为ScatterDimensionNumbers并调用lax.scatter*系列原语(scatter.py#L43-L90)。这也解释了为何jax.ops.segment_*与.at[...].add()共享同一套底层机制:二者本质都是 scatter。
5.2 布尔类型的段归约
测试 tests/lax_numpy_indexing_test.py#L1843-L1856(testSegmentReduceBoolean)覆盖了segment_min/segment_max在bool_类型下的行为,恒等元分别为True与False(对应布尔“与”与“或”的天然语义),并组合测试了bucket_size=[None, 2]、num_segments=[None, 1, 3]等参数组合。
5.3 形状多态(Shape Polymorphism)
段归约同样出现在形状多态测试中: tests/shape_poly_test.py#L3464-L3467 将四个函数与("max", ops.segment_max)、("min", ops.segment_min)、("sum", ops.segment_sum)、("prod", ops.segment_prod)一一配对做多态维度测试。这印证了segment_*系列在动态形状场景(如jax.jit配合形状抽象)下也是可用的——但前提仍是num_segments保持静态。
六、实战要点与限制小结
- 迁移遗留代码:仓库中任何
jax.ops.index_update/index_add等调用都应改写为x.at[idx].set(y)/x.at[idx].add(y)形式;当前jax.ops仅导出四个segment_*函数(见 jax/ops/init.py)。 - JIT 中使用必须显式
num_segments:否则concrete_dim_or_error会拒绝编译;推荐把num_segments作为static_argnums传入。 - 段 id 无需有序、可以重复:默认模式(
FILL_OR_DROP)下越界 id 被丢弃;负 id 默认按 Python 语义从尾部计数。 - 可微性注意:
segment_sum全程可微;segment_prod/segment_min/segment_max在unique_indices=False(存在重复段 id)时梯度未实现,会在反向传播时报NotImplementedError。 - 追求数值稳定性:sum/prod 类长段归约可考虑
bucket_size分桶;追求吞吐时可通过indices_are_sorted=True、unique_indices=True换取后端更优执行路径。 - 分布式场景:通过
out_sharding(配合jax.jit与 mesh)可为输出指定分片;注意其与bucket_size互斥、prod/min/max尚不支持unreduced分片规格。
延伸阅读:索引更新的完整模式语义参见jax.lax.GatherScatterMode(源码 jax/_src/lax/slicing.py);segment_*全部实现与文档字符串见 jax/_src/ops/scatter.py;行为验证可运行 tests/lax_numpy_indexing_test.py 中的testSegmentSum等用例。
【免费下载链接】jaxComposable transformations of Python+NumPy programs: differentiate, vectorize, JIT to GPU/TPU, and more项目地址: https://gitcode.com/GitHub_Trending/ja/jax
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考