JAX 段归约操作指南:jax.ops.segment_sum / segment_max 系列函数与 .at 索引更新全面解析
JAX 段归约操作指南jax.ops.segment_sum / segment_max 系列函数与 .at 索引更新全面解析【免费下载链接】jaxComposable transformations of PythonNumPy 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 模块定位与公开 APIjax.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 functionsjax.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] yx x.at[idx].add(y)x[idx] yx x.at[idx].subtract(y)x[idx] - yx x.at[idx].multiply(y)x[idx] * yx x.at[idx].divide(y)x[idx] / yx x.at[idx].power(y)x[idx] ** yx 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()中的越界索引被裁剪clipset()/add()等更新中的越界索引被丢弃drop。clip将越界索引裁剪到合法范围。drop忽略越界索引。filldrop的别名对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.], dtypefloat32) x.at[10].add(10) # 默认模式越界更新被丢弃 Array([0., 1., 2., 3., 4.], dtypefloat32) x.at[20].add(10, modeclip) # clip 模式裁剪到末尾 Array([ 0., 1., 2., 3., 14.], dtypefloat32) x.at[20].get() # get() 默认裁剪 Array(4., dtypefloat32) x.at[20].get(modefill, fill_value-1) Array(-1., dtypefloat32)在实现层面.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_segmentsNone, indices_are_sortedFalse, unique_indicesFalse, bucket_sizeNone, modeNone, out_shardingNone) segment_prod(data, segment_ids, num_segmentsNone, indices_are_sortedFalse, unique_indicesFalse, bucket_sizeNone, modeNone, out_shardingNone) segment_max(data, segment_ids, num_segmentsNone, indices_are_sortedFalse, unique_indicesFalse, bucket_sizeNone, modeNone, out_shardingNone) segment_min(data, segment_ids, num_segmentsNone, indices_are_sortedFalse, unique_indicesFalse, bucket_sizeNone, modeNone, out_shardingNone)3.1 基础示例摘自源码 docstringfrom 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], dtypeint32) # [01, 23, 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], dtypeint32) # [0*1, 2*3, 4*5] # --- segment_max段最大值 --- segment_max(data, segment_ids) # Array([1, 3, 5], dtypeint32) # --- segment_min段最小值 --- segment_min(data, segment_ids) # Array([0, 2, 4], dtypeint32)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]], dtypeint32)四、参数语义深度解析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_argnums2)(data, segment_ids, 3) # Array([1, 5, 4], dtypeint32)4.2mode越界段 id 的处理默认modeNone实际映射为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_segments4) # Array([5, 2, 3, 3]) # 越界 id4、8与规范化后的负 id 按 mod 折回/丢弃4.3indices_are_sorted与unique_indices性能提示indices_are_sortedTrue声明segment_ids规范化后升序排列。若声明与实际不符输出未定义。unique_indicesTrue声明每个段 id 至多出现一次。若声明与实际不符输出未定义。这两项本质上是把“保证”交给用户、换取部分后端更高效的执行路径源码在 scatter.py#L145-L147 将其并入底层 scatter 调用。同时注意segment_prod/segment_max/segment_min的自动微分只在unique_indicesTrue时完整实现——例如测试 tests/lax_numpy_indexing_test.py#L1782-L1785 断言scatter_mul的梯度在非唯一索引下会抛出NotImplementedError“scatter_mul gradients are only implemented ifunique_indicesTrue”。4.4bucket_size数值稳定性分桶默认None表示不分桶。传入正整数时算法会把segment_ids按顺序切成多个桶每桶最多bucket_size个元素在每个桶内单独执行段归约再对桶间结果做二次归约reducer(out, axis0)。这样做的动机是改善 sum/prod 这类累积归约的数值稳定性源码注释 scatter.py#L205-L206Bucketize 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_axesscatter.py#L86-L90。两点实现限制值得注意传入out_sharding时不能再同时使用bucket_sizescatter.py#L208-L209 直接raise NotImplementedErrorsegment_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_shardingout_s并配合jax.jit使用。五、底层实现原理一切归约为 scatter5.1 共享的_segment_update流水线四个函数都是薄封装最终汇聚到同一个私有函数_segment_updatescatter.py#L175-L218区别仅在于底层 scatter 原语与归约器公开函数底层 scatter 原语归约器恒等元identitysegment_sumscatter_addsum0segment_prodscatter_mulprod1segment_minscatter_minmininf整数为类型最大值布尔为Truesegment_maxscatter_maxmax-inf整数为类型最小值布尔为False执行流程为校验输入check_arraylike、将data与segment_ids转为 JAX 数组确定num_segments默认max(segment_ids)1并做静态性校验用_get_identity(op, dtype)求得该归约的恒等元scatter.py#L153-L172scatter_min对布尔取True、整数取iinfo(dtype).max、浮点取infscatter_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-L1856testSegmentReduceBoolean覆盖了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_indicesFalse存在重复段 id时梯度未实现会在反向传播时报NotImplementedError。追求数值稳定性sum/prod 类长段归约可考虑bucket_size分桶追求吞吐时可通过indices_are_sortedTrue、unique_indicesTrue换取后端更优执行路径。分布式场景通过out_sharding配合jax.jit与 mesh可为输出指定分片注意其与bucket_size互斥、prod/min/max尚不支持unreduced分片规格。延伸阅读索引更新的完整模式语义参见jax.lax.GatherScatterMode源码 jax/_src/lax/slicing.pysegment_*全部实现与文档字符串见 jax/_src/ops/scatter.py行为验证可运行 tests/lax_numpy_indexing_test.py 中的testSegmentSum等用例。【免费下载链接】jaxComposable transformations of PythonNumPy programs: differentiate, vectorize, JIT to GPU/TPU, and more项目地址: https://gitcode.com/GitHub_Trending/ja/jax创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考