Ray Data 公共 API 全景:基于 autosummary 索引解析 Dataset、DataIterator 与聚合体系
【免费下载链接】rayRay is an AI compute engine. Ray consists of a core distributed runtime and a set of AI Libraries for accelerating ML workloads.项目地址: https://gitcode.com/gh_mirrors/ra/ray
Ray Data 是 Ray 面向 AI 工作负载的分布式数据处理引擎,围绕 Arrow 块(Block)构建了一套"延迟变换 + 流式执行"的数据集抽象。本文以仓库中 doc/source/data/api/_autogen.rst 这份自动生成 API 文档的索引清单为骨架,逐一解析其中列出的 8 个核心公共 API(Dataset、DataIterator、Schema、DatasetSummary、GroupedData、AggregateFn、AggregateFnV2),并结合 python/ray/data 下的源码实现与同目录 API 参考文档,讲清每个类的职责边界、典型用法与底层原理。读完本文,你将能够:看懂 Ray Data API 文档的组织与生成机制,掌握数据集从创建、变换到消费的完整调用链,并具备编写自定义分布式聚合(AggregateFnV2)的实战能力。
一、_autogen.rst是什么:Ray Data API 文档的生成枢纽
_autogen.rst是一份 Sphinx 的autosummary指令文件,文件头注释明确说明了它的用途:仅用于自动生成 API 文档,不应被纳入 toctree。任何需要出现在 API 文档中的类,都需要先加入这份清单,随后 Sphinx 会为清单中的每个条目生成对应的ray.data.<class>.rst文件,再由顶层 RST 文件包含这些生成产物。
该文件的完整清单如下:
.. currentmodule:: ray.data .. autosummary:: :nosignatures: :template: autosummary/class_v2.rst :toctree: DataIterator Dataset Schema stats.DatasetSummary grouped_data.GroupedData aggregate.AggregateFn aggregate.AggregateFnV2几个值得注意的细节:
.. currentmodule:: ray.data将后续条目的默认模块设置为ray.data,因此清单中的Schema实际指向ray.data.Schema,而带模块前缀的stats.DatasetSummary、grouped_data.GroupedData、aggregate.AggregateFn、aggregate.AggregateFnV2则分别指向ray.data.stats、ray.data.grouped_data、ray.data.aggregate三个子模块。:nosignatures:表示生成条目时不展示函数签名,适合以类为单位的索引页;:template: autosummary/class_v2.rst指定了类模板;:toctree:让生成的 RST 文件落入 toctree,形成可导航的子页面。- 这 7 个条目(含 8 个类)恰好构成了 Ray Data 用户最常接触的"门面 API":数据集本体(
Dataset)、迭代消费接口(DataIterator)、类型描述(Schema)、统计汇总(DatasetSummary)、分组中间态(GroupedData)以及两代自定义聚合基类(AggregateFn/AggregateFnV2)。
从源码结构看,这一清单与 python/ray/data 模块根目录下的模块划分一一对应:dataset.py定义Dataset与Schema,iterator.py定义DataIterator,grouped_data.py定义GroupedData,aggregate.py定义聚合基类,stats.py定义DatasetSummary。整个ray.data包还有read_api.py、block.py、datatype.py、context.py等模块,共同支撑起完整的数据集能力。
二、核心数据抽象:Dataset—— 分布式数据集合
Dataset是 Ray Data 的基石。源码 python/ray/data/dataset.py 给出了精确定义:
A Dataset is a distributed data collection for data loading and processing. Datasets are distributed pipelines that produce
ObjectRef[Block]outputs, where each block holds data in Arrow format, representing a shard of the overall data collection. The block also determines the unit of parallelism.
关键信息可以拆解为三点:
- 块(Block)即并行单元:一个 Dataset 在物理上被切分为若干 Arrow 格式的块,每块对应一个 Ray
ObjectRef,块的数量决定了并行度。这意味着用户看到的是一个逻辑上的数据集,底层却是分布式对象存储中的一组可并行处理的分片。 - 三种创建途径:从外部存储(本地磁盘、S3、HDFS 等)通过
read_*()API 读入;从内存数据通过from_*()API 构造;从合成数据通过range_*()API 生成。例如:
import ray # 合成数据 ds = ray.data.range(1000) # 内存数据 ds = ray.data.from_items([{"col1": i, "col2": i * 2} for i in range(1000)]) # 外部存储(对象存储/文件系统) ds = ray.data.read_parquet("s3://bucket/path") # 写回外部存储 ds.write_csv("s3://bucket/output")- 变换与消费分离、执行惰性化:
Dataset有两类操作——变换(transformation)产生新 Dataset,如map_batches()、sort()、random_shuffle()、repartition()、join();消费(consumption)产生具体值,如iter_batches()、take_all()、max()。变换是惰性的,只有下游消费触发时才真正执行,这为执行引擎做算子融合与流式调度留出了空间:
>>> ds = ray.data.range(1000) >>> ds.max("id") # 消费类操作,立即触发执行 999 >>> ds.random_shuffle() # 变换类操作,返回新的 Dataset(未物化) shape: (1000, 1) ╭───────╮ │ id │ │ --- │ │ int64 │ ╰───────╯ (Dataset isn't materialized)在 doc/source/data/api/dataset.rst 中,Dataset的 API 参考还延伸出三组周边 API:
- Compute Strategy(计算策略):
ActorPoolStrategy与TaskPoolStrategy,用于控制map_batches等变换采用 Actor 池(可复用状态、适合 GPU 与重量级初始化)还是轻量 Task(适合无状态短任务)来执行。 - Schema 与 DatasetSummary:见本文第四、五节。
- Developer API 与 Deprecated API:
to_pandas_refs、to_numpy_refs、to_arrow_refs、iter_internal_ref_bundles等内部能力,以及to_random_access_dataset、iter_tf_batches等已废弃接口,后者仍保留在文档中以帮助用户平滑迁移。
三、消费侧抽象:DataIterator—— 单次遍历读取全部记录
DataIterator是定义在 python/ray/data/iterator.py 的抽象基类(abc.ABC),用于从Dataset读取记录。它的核心语义(源码 docstring)是:
For Datasets, each iteration call represents a complete read of all items in the Dataset.
即每次对DataIterator的迭代都是一次对 Dataset 全量数据的完整读取。Dataset.iterator()可获取一个迭代器:
>>> ds = ray.data.range(5) >>> ds.iterator() DataIterator(shape: (5, 1) ╭───────╮ │ id │ │ --- │ │ int64 │ ╰───────╯ (Dataset isn't materialized))DataIterator是 Ray Data 与下游训练框架衔接的关键接口:
- 在 Ray Train 中,每个 Trainer Actor 应当通过
ray.train.get_dataset_shard("train")获取属于自己的迭代器分片,实现数据在分布式训练 worker 间的切分; - 迭代器提供
iter_batches()、iter_torch_batches()、iter_tf_batches()等批量读取方法,可直接产出 PyTorch / TensorFlow 所需的张量批次,是"数据处理管线 → 训练循环"之间的标准桥梁。
底层实现类为ray.data._internal.iterator.iterator_impl.DataIteratorImpl,它持有_to_ref_bundle_iterator()等内部方法,将迭代操作映射到底层的 RefBundle 流式执行之上(见 python/ray/data/_internal/iterator/iterator_impl.py),相关参考见 doc/source/data/api/data_iterator.rst。
四、类型与统计:Schema与DatasetSummary
4.1Schema:数据集的列类型描述
Schema定义在 python/ray/data/dataset.py,是对底层 PyArrow Schema(或 Pandas Block Schema)的一层封装:
base_schema属性保存底层 Arrow / Pandas schema;names属性返回用户可见的列名列表(实现中会过滤掉__bsp_stub这类读路径注入的物理占位列);types属性返回各列的 Arrow 类型,对非 Arrow 兼容类型统一返回object;实现中针对pd.ArrowDtype、pd.StringDtype、掩码类型(BaseMaskedDtype)做了到pyarrow.DataType的转换,并将 Ray 的张量扩展类型(create_arrow_fixed_shape_tensor_type)纳入类型系统。
一个值得注意的实现细节:Schema在构造时会快照当前的DataContext(copy.deepcopy(DataContext.get_current())),注释表明"数据集创建时刻的配置决定了其行为",即 Dataset 的配置语义在创建时被冻结,避免运行期上下文漂移。
4.2DatasetSummary:数据集统计汇总
DatasetSummary定义在 python/ray/data/stats.py,是一个标注了stability="alpha"的@dataclass,用于承载对数据集列计算出的统计信息(如 min/max/count/missing_pct 等)。
其实现有一个值得称道的细节:由于聚合结果可能与原列类型不一致(例如字符串列的count是 int64),DatasetSummary内部把统计拆成两张 PyArrow 表——_stats_matching_column_dtype(与原列同类型的统计,保留原 dtype)和_stats_mismatching_column_dtype(如 count、missing_pct 等类型不同的统计),并在to_pandas()时合并为单个 DataFrame;遇到TensorExtensionType等扩展类型转换失败时,会降级为逐列转换、将问题列转为 null 类型(见 python/ray/data/stats.py)。这保证了统计结果既能保持类型严谨,又对 pandas 用户足够友好。
五、分组与聚合:GroupedData与两代AggregateFn
5.1GroupedData:惰性分组中间态
GroupedData定义在 python/ray/data/grouped_data.py,docstring 明确指出:
Represents a grouped dataset created by calling
Dataset.groupby(). The actual groupby is deferred until an aggregation is applied.
也就是说,ds.groupby("key")并不会立即执行任何分组逻辑,只是返回一个持有dataset、key、num_partitions的轻量句柄;真正的分组发生在调用aggregate()之时。aggregate(*aggs)将分组键与聚合算子打包成一个Aggregate逻辑算子,接入 Dataset 的逻辑计划(LogicalPlan),并返回新的Dataset。结果数据集有n + 1列:第一列是分组键,其后依次是各聚合结果;当分组键为None时键列省略。
GroupedData还提供map_groups()方法,支持对每个组应用用户自定义的批处理函数(fn),并接受compute、batch_format、num_cpus、num_gpus、concurrency等资源与并行度参数,适合"分组后做复杂变换"的场景。
5.2 内置聚合函数
在 doc/source/data/api/aggregate.rst 中,聚合 API 共列出 18 个条目:基类AggregateFnV2、AggregateFn与 16 个内置实现:
Count、Sum、Min、Max、Mean、Std、AbsMax、Quantile、Unique、AsList、CountDistinct、ValueCounter、MissingValuePercentage、ZeroPercentage、ApproximateQuantile、ApproximateTopK。
这些内置聚合统一通过Dataset.aggregate()或Dataset.groupby().aggregate()使用。例如:
import ray from ray.data.aggregate import Mean, ApproximateTopK ds = ray.data.from_items( [{"group": "A", "score": 1.0}, {"group": "A", "score": 3.0}, {"group": "B", "score": 2.0}] ) # 全量聚合 print(ds.aggregate(Mean(on="score")).take_all()) # 分组聚合 print(ds.groupby("group").aggregate(Mean(on="score")).take_all())5.3AggregateFn(已废弃):基于 init/merge/accumulate 的自定义聚合
AggregateFn定义在 python/ray/data/aggregate.py,源码中已被标记@Deprecated,提示"请使用AggregateFnV2代替",但它定义了自定义聚合最经典的四段式接口,理解它有助于理解 V2:
init(key):为每个组创建初始累加器(如 0、空列表或空字典);accumulate_row/accumulate_block:逐行或逐块更新累加器(两者必须且只能提供一个,支持向量化批量处理);merge(c1, c2):合并不同 worker 产生的两个累加器;finalize(acc):可选,将最终累加器转换为输出(不提供则原样返回)。
经典示例(按组计数):
import ray from ray.data.aggregate import AggregateFn count_agg = AggregateFn( init=lambda k: 0, accumulate_row=lambda counter, row: counter + 1, merge=lambda c1, c2: c1 + c2, name="custom_count", ) ds = ray.data.from_items([{"group": "A"}, {"group": "B"}, {"group": "A"}]) result = ds.groupby("group").aggregate(count_agg).take_all() # result: [{'group': 'A', 'custom_count': 2}, {'group': 'B', 'custom_count': 1}]5.4AggregateFnV2:当前推荐的高效聚合接口
AggregateFnV2定义在 python/ray/data/aggregate.py,继承自AggregateFn并参数化为Generic[AccumulatorType, AggOutputType],是当前官方推荐的自定义聚合方式。它的执行被明确规范为四步:
- 初始化(Initialization):对每个分组(或整个数据集)用
zero_factory创建初始累加器; - 块聚合(Block Aggregation):
aggregate_block独立作用于每个块,产出该块的局部聚合结果(可向量化,性能关键); - 合并(Combination):
combine将各块的局部结果合并为统一累加器; - 终结(Finalization):可选的
finalize将最终累加器转换为输出格式。
构造参数为name(输出列名,如"sum(my_col)")、zero_factory(初始零值工厂,如求和用lambda: 0、求最小用lambda: float("inf"))、on(聚合列名,None表示整行,如Count())与ignore_nulls(是否跳过 null)。
两个泛型参数的含义在源码中有清晰示例:
Count(AggregateFnV2[int, int]) # 累加器: int, 输出: int Sum(AggregateFnV2[Union[int, float], ...]) # 累加器: 数值, 输出: 数值 Mean(AggregateFnV2[List[...], float]) # 累加器: [sum, count], 输出: float Std(AggregateFnV2[List[...], float]) # 累加器: [M2, mean, count], 输出: float即AccumulatorType是中间状态类型(aggregate_block的产出、combine的输入输出、finalize的入参),AggOutputType是写回结果集的最终类型。Mean用[sum, count]作为累加器再算出均值、Std用[M2, mean, count]承载二阶矩——这种"复合累加器"正是 V2 接口表达力的体现。一个简单的自定义Sum子类骨架如下:
from typing import Union from ray.data.aggregate import AggregateFnV2 class MySum(AggregateFnV2[Union[int, float], Union[int, float]]): def __init__(self, on: str): super().__init__( name=f"sum({on})", zero_factory=lambda: 0, on=on, ignore_nulls=True, ) def aggregate_block(self, block, column_name=None, *args, **kwargs): # 对单个块做向量化求和,返回局部累加器 ... def combine(self, left_acc, right_acc): return left_acc + right_acc # finalize 可省略:默认原样返回累加器作为输出六、数据加载与保存 API 全景
围绕上述核心类,Ray Data 提供了极广的数据进出通道,分别收录在 doc/source/data/api/loading_data.rst 与 doc/source/data/api/saving_data.rst 中,这也是 API 参考页的主体。
6.1 加载侧(Loading Data API)
公共 API 按数据源组织,主要包括:
- 内存/框架互转:
from_pandas、from_numpy、from_arrow、from_items、from_torch、from_tf、from_dask、from_modin、from_mars、from_daft、from_spark、from_huggingface; - 文件与列式格式:
read_parquet、read_csv、read_json、read_orc、read_avro、read_numpy、read_text、read_tfrecords、read_mcap、read_webdataset、read_zarr; - AI 多模态数据:
read_images、read_videos、read_audio、read_lerobot、read_binary_files; - 湖仓与数据库:
read_bigquery、read_clickhouse、read_mongo、read_sql、read_snowflake、read_delta、read_delta_sharing_tables、read_hudi、read_iceberg、read_lance、read_kafka、read_databricks_tables、read_unity_catalog(含Catalog/DatabricksUnityCatalog/ReaderFormat类); - 合成数据:
range、range_tensor; - 分区与重排辅助:
datasource.Partitioning、datasource.PartitionStyle、datasource.PathPartitionFilter、datasource.PathPartitionParser、FileShuffleConfig; - 开发者 API:
from_arrow_refs、from_numpy_refs、from_pandas_refs、Datasource、datasource.FileBasedDatasource、read_datasource、ReadTask,以及元数据提供者datasource.BaseFileMetadataProvider、datasource.DefaultFileMetadataProvider、datasource.FileMetadataProvider。
6.2 保存侧(Saving Data API)
保存侧与加载侧对称,公共 API 包括Dataset.write_parquet、write_csv、write_json、write_orc、write_numpy、write_tfrecords、write_images、write_lance、write_mongo、write_clickhouse、write_bigquery、write_snowflake、write_sql、write_iceberg,以及框架互转to_pandas、to_numpy、to_dask、to_modin、to_mars、to_daft、to_spark;开发者侧还有Datasink、Dataset.write_datasink、datasource.RowBasedFileDatasink、datasource.BlockBasedFileDatasink、datasource.WriteResult、datasource.WriteReturnType、datasource.FilenameProvider等扩展点。这意味着用户既可以开箱即用地落盘,也可以基于 Datasink 定制自己的输出格式。
七、扩展阅读与配套 API
- 执行配置:
ExecutionOptions与ExecutionResources(见 doc/source/data/api/execution_options.rst),用于控制数据集执行的资源上限与策略; - 数据类型:
DataType与TypeCategory(见 doc/source/data/api/datatype.rst),描述数据集支持的数据类型分类; - Checkpoint / 表达式 / DataContext:分别见 checkpoint.rst、expressions.rst、data_context.rst;preprocessor.rst 与 llm.rst 覆盖预处理与 LLM 生态集成;
- API 总入口:doc/source/data/api/api.md 通过 toctree 将上述所有 API 页面组织为完整的 Ray Data API 章节,
_autogen.rst生成的类文档即嵌入其中。
整体上,Ray Data 的 API 设计遵循"小而稳的门面 + 可扩展的内部接口"原则:Dataset、DataIterator、Schema等门面类保持稳定,而AggregateFnV2、Datasource、Datasink、BlockAccessor等接口则向开发者开放深度定制能力。无论是快速做数据预处理,还是在训练管线中切分数据批次,或是在数千核集群上实现自定义分布式聚合,这套 API 都提供了从入口到扩展的完整路径。
【免费下载链接】rayRay is an AI compute engine. Ray consists of a core distributed runtime and a set of AI Libraries for accelerating ML workloads.项目地址: https://gitcode.com/gh_mirrors/ra/ray
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考