news 2026/9/11 18:19:12

JAX 段归约操作指南:jax.ops.segment_sum / segment_max 系列函数与 .at 索引更新全面解析

作者头像

张小明

前端开发工程师

1.2k 24
文章封面图
JAX 段归约操作指南:jax.ops.segment_sum / segment_max 系列函数与 .at 索引更新全面解析

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_sumsegment_prodsegment_minsegment_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 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_updatejax.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 的关键差异:

  1. 纯函数性:任何x.at[...]表达式都不会修改原数组x,而是返回修改后的副本;不过在jax.jit编译的函数内部,x = x.at[idx].set(y)这类表达式保证会被原地应用(buffer donation 优化)。
  2. 重复索引语义:与 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_sortedunique_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.scatterscatter_addscatter_mulscatter_minscatter_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_sortedunique_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参数(接受NamedShardingPartitionSpec),用于指定输出的分片布局。源码先通过canonicalize_sharding(out_sharding, 'segment_xxx')规范化,再沿网格显式轴调用auto_axes(scatter.py#L86-L90)。两点实现限制值得注意:

  • 传入out_sharding不能再同时使用bucket_size(scatter.py#L208-L209 直接raise NotImplementedError);
  • segment_prodsegment_maxsegment_minunreduced(未归约)分片规格尚未支持(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_sumscatter_addsum0
segment_prodscatter_mulprod1
segment_minscatter_minmin+inf(整数为类型最大值,布尔为True
segment_maxscatter_maxmax-inf(整数为类型最小值,布尔为False

执行流程为:

  1. 校验输入(check_arraylike)、将datasegment_ids转为 JAX 数组;
  2. 确定num_segments(默认max(segment_ids)+1)并做静态性校验;
  3. _get_identity(op, dtype)求得该归约的恒等元(scatter.py#L153-L172:scatter_min对布尔取True、整数取iinfo(dtype).max、浮点取infscatter_max反之);
  4. 以恒等元构造形状为(num_segments,) + data.shape[1:]的输出缓冲区out
  5. 调用_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_maxbool_类型下的行为,恒等元分别为TrueFalse(对应布尔“与”与“或”的天然语义),并组合测试了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保持静态。

六、实战要点与限制小结

  1. 迁移遗留代码:仓库中任何jax.ops.index_update/index_add等调用都应改写为x.at[idx].set(y)/x.at[idx].add(y)形式;当前jax.ops仅导出四个segment_*函数(见 jax/ops/init.py)。
  2. JIT 中使用必须显式num_segments:否则concrete_dim_or_error会拒绝编译;推荐把num_segments作为static_argnums传入。
  3. 段 id 无需有序、可以重复:默认模式(FILL_OR_DROP)下越界 id 被丢弃;负 id 默认按 Python 语义从尾部计数。
  4. 可微性注意segment_sum全程可微;segment_prod/segment_min/segment_maxunique_indices=False(存在重复段 id)时梯度未实现,会在反向传播时报NotImplementedError
  5. 追求数值稳定性:sum/prod 类长段归约可考虑bucket_size分桶;追求吞吐时可通过indices_are_sorted=Trueunique_indices=True换取后端更优执行路径。
  6. 分布式场景:通过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),仅供参考

版权声明: 本文来自互联网用户投稿,该文观点仅代表作者本人,不代表本站立场。本站仅提供信息存储空间服务,不拥有所有权,不承担相关法律责任。如若内容造成侵权/违法违规/事实不符,请联系邮箱:809451989@qq.com进行投诉反馈,一经查实,立即删除!
网站建设 2026/9/11 18:16:47

【Rust入门知识点学与练】第17课:闭包 Closures

知识点1&#xff1a;闭包基本语法 闭包是能捕获周围环境中变量的匿名函数&#xff1a; fn main() {// 完整写法let add_one |x: i32| -> i32 { x 1 };println!("{}", add_one(5)); // 6// 简写&#xff1a;类型可以省略&#xff08;编译器会推断&#xff09;le…

作者头像 李华
网站建设 2026/9/11 18:16:27

顶空气体分析与残氧仪技术详解及应用

1. 顶空气体分析技术概述 在现代包装工业中&#xff0c;产品保质期的延长和品质保持是核心诉求。顶空气体分析技术&#xff08;Headspace Gas Analysis&#xff09;作为一种非破坏性检测方法&#xff0c;通过分析包装内部顶部空间的气体成分&#xff0c;为包装工艺优化和产品质…

作者头像 李华
网站建设 2026/9/11 18:16:01

垃圾短信识别实战:从TF-IDF到多模型对比的课程设计全解析

简介&#xff1a;这是一套面向网络数据挖掘课程设计/实训的垃圾短信识别系统完整工程&#xff0c;适合高校学生在毕业设计、课程设计、大作业、工程实训或学科竞赛中直接复现与二次开发。项目基于机器学习与自然语言处理实现短信文本分类&#xff0c;从数据预处理、特征工程到S…

作者头像 李华
网站建设 2026/9/11 18:14:50

如何三步跑通离线语音识别:Vosk 50MB 轻量模型完整指南

如何三步跑通离线语音识别&#xff1a;Vosk 50MB 轻量模型完整指南 【免费下载链接】vosk-api Offline speech recognition API for Android, iOS, Raspberry Pi and servers with Python, Java, C# and Node 项目地址: https://gitcode.com/GitHub_Trending/vo/vosk-api …

作者头像 李华
网站建设 2026/9/11 18:13:33

iphreeqc-py 0.1a6 安装与使用:用 Python 驱动 PHREEQC 水化学模拟

简介&#xff1a;iphreeqc-py 是一个面向 Python 开发者的第三方库安装包&#xff0c;专为地球化学模拟、水质分析以及相关科学计算场景设计。该库来自官方渠道&#xff0c;本质上是 PHREEQC 经典水文地球化学模拟引擎的 Python 接口封装&#xff0c;用户可以在熟悉的 Python 环…

作者头像 李华
网站建设 2026/9/11 18:12:14

燃气用户管理系统有哪些?核心功能与选型指南

燃气是城市能源供应体系的重要组成部分&#xff0c;覆盖居民生活与工业生产两大场景。随着管道燃气用户规模持续扩大&#xff0c;燃气企业既要在经营端完成数以万计的用户建档、计费与缴费服务&#xff0c;又要在安全端落实入户安检、隐患整改等主体责任&#xff0c;传统手工台…

作者头像 李华