Ray Data GroupedData API 深度指南:基于 Dataset.groupby() 的分组聚合与 map_groups 实践
【免费下载链接】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
GroupedData 是 Ray Data 中由Dataset.groupby()返回的懒加载分组对象,它把"分组"这一操作推迟到真正执行聚合时才发生。本文以 doc/source/data/api/grouped_data.rst 为核心骨架,结合 GroupedData 实现源码 与 e2e 测试用例,系统讲解分组键的指定方式、七类内置聚合(count/sum/min/max/mean/std/aggregate)、自定义聚合器(AggregateFn / AggregateFnV2)以及 map_groups 分组变换函数的完整用法与底层执行原理,帮助你用最少的代码完成按列统计、分组归一化、组内特征变换等常见数据处理任务。
一、GroupedData 是什么:groupby 返回的懒加载分组对象
在 Ray Data 中,分组操作遵循一个"两步走"的模式:
- 先调用
Dataset.groupby(key)得到一个GroupedData对象; - 再在该对象上调用聚合方法(如
.sum()、.mean())或变换方法(如.map_groups())触发实际计算。
官方 API 文档 grouped_data.rst 明确指出:
The groupby call returns GroupedData objects:
Dataset.groupby()。
对应的类定义位于 python/ray/data/grouped_data.py#L29-L33,其文档字符串对"懒加载"语义做了最精炼的概括:
class GroupedData: """Represents a grouped dataset created by calling ``Dataset.groupby()``. The actual groupby is deferred until an aggregation is applied. """也就是说,groupby()本身不触发任何数据搬运或排序,它只是记录三个关键状态(见 构造函数):
| 内部属性 | 含义 |
|---|---|
_dataset | 被分组的原始Dataset |
_key | 分组键:单个列名(str)、列名列表(List[str]),或None(表示全部行归入一个全局组) |
_num_partitions | 仅在使用哈希洗牌策略时生效的目标分区数,未设置时默认取DataContext.min_parallelism |
GroupedData的构造函数是内部 API(源码注释明确 "The constructor is not part of the GroupedData API"),用户唯一正确的创建方式就是调用Dataset.groupby()。
二、创建分组:Dataset.groupby() 的参数与语义
groupby()定义在 python/ray/data/dataset.py#L3782-L3840,签名如下:
def groupby( self, key: Union[str, List[str], None], num_partitions: Optional[int] = None, ) -> "GroupedData":2.1 key:分组键的三种形态
- 单个列名(str):按某一列的值分组,如
ds.groupby("variety"); - 列名列表(List[str]):按多列的组合值分组,例如
ds.groupby(["year", "month"]); None:所有行归入一个全局组,等价于"对整个数据集做一次聚合"。
从实现看,groupby()在key非空时会调用SortKey(key).validate_schema(...)校验键列确实存在于数据集 schema 中(见 dataset.py#L3832-L3835);对key=None则总是放行,因为None被解释为"全局单组"。同时,num_partitions <= 0会直接抛出ValueError("num_partitionsmust be a positive integer")。
import ray ds = ray.data.from_items([ {"group": 1, "value": 1}, {"group": 1, "value": 2}, {"group": 2, "value": 3}, {"group": 2, "value": 4}, ]) # 按单列分组 grouped = ds.groupby("group") # 按多列分组 grouped_multi = ds.groupby(["group", "value"]) # 全局单组(对整个数据集聚合) global_grouped = ds.groupby(None)2.2 num_partitions:控制哈希洗牌的目标分区数
文档与源码都强调num_partitions仅在采用哈希洗牌(hash shuffle)策略时才相关。它在 GroupedData.map_groups 的实现 中被使用:
if self._key is None: shuffled_ds = self._dataset.repartition(1) elif self._dataset.context.shuffle_strategy in ( ShuffleStrategy.HASH_SHUFFLE, ShuffleStrategy.SHUFFLE_V2, ShuffleStrategy.GPU_SHUFFLE, ): num_partitions = ( self._num_partitions or self._dataset.context.default_hash_shuffle_parallelism ) shuffled_ds = self._dataset.repartition( num_partitions, keys=self._key, sort=True ) else: shuffled_ds = self._dataset.sort(self._key)这段代码揭示了三种不同的分组预处理路径:
| 场景 | 底层操作 | 说明 |
|---|---|---|
key is None | repartition(1) | 全部数据收敛到 1 个 block,形成单一全局组 |
| 使用哈希/新式/GPU 洗牌策略 | repartition(num_partitions, keys=key, sort=True) | 按键哈希分区,并在分区内排序,保证相同键值的行被共置(co-located) |
| 其他策略(默认排序式洗牌) | sort(key) | 对数据集做全局排序,让相同键值的行相邻 |
因此groupby()的时间复杂度约为O(dataset size * log(dataset size / parallelism))(见 groupby 文档字符串)。
三、内置聚合:count / sum / min / max / mean / std
GroupedData 提供了 6 个开箱即用的统计聚合方法,全部属于CDS_API_GROUP("Computations or Descriptive Stats")API 组。它们共享同一套on与ignore_nulls语义,底层都通过_aggregate_on()帮助函数(见 grouped_data.py#L79-L96)转换为对aggregate()的多聚合调用。
3.1 方法签名与参数语义
def count(self) -> Dataset: ... def sum(self, on: Union[str, List[str]] = None, ignore_nulls: bool = True) -> Dataset: ... def min(self, on: Union[str, List[str]] = None, ignore_nulls: bool = True) -> Dataset: ... def max(self, on: Union[str, List[str]] = None, ignore_nulls: bool = True) -> Dataset: ... def mean(self, on: Union[str, List[str]] = None, ignore_nulls: bool = True) -> Dataset: ... def std(self, on: Union[str, List[str]] = None, ddof: int = 1, ignore_nulls: bool = True) -> Dataset: ...on参数(sum/min/max/mean/std)有三种取值,返回结构随之变化(源码在 grouped_data.py#L431-L441 等处有完整说明):
on=None:对数据集中除分组键外的每一列分别做聚合,返回"分组键列 + 每列一个聚合结果列";on="col":只对指定列聚合,返回两列(键列 + 结果列);on=["col_1", ..., "col_n"]:对多个列分别聚合,返回n + 1列(键列 + n 个结果列)。
ignore_nulls参数:True(默认)时忽略空值计算;False时只要遇到空值(np.nan、None、pd.NaT均视为空值)输出即为空。
ddof参数(仅 std):Delta Degrees of Freedom,除数采用N - ddof,默认ddof=1(样本标准差),ddof=0对应总体标准差。
全局组特例:当分组键为None时,所有返回结果都会省略键列,只输出聚合列。
3.2 完整示例
import ray # 构造 100 条记录,A 为 0/1/2 三类,B、C 为待聚合数值 ds = ray.data.from_items( [{"A": i % 3, "B": i, "C": i ** 2} for i in range(100)] ) # 按 A 分组,对 B、C 分别求和(on 传列名列表) ds.groupby("A").sum(["B", "C"]) # 按 A 分组,对 B 求均值 ds.groupby("A").mean("B") # 按 A 分组,求每组行数,返回 [k, v] 两列 ds.groupby("A").count() # 按 A 分组,同时求 B、C 的最小值 ds.groupby("A").min(["B", "C"]) # 按 A 分组,对 B 求总体标准差(ddof=0) ds.groupby("A").std("B", ddof=0) # 全局分组:对整个数据集做列级求和(无键列) ds.groupby(None).sum()count()的返回结构是[k, v]两列,其中k为分组键值、v为该键的行数;key=None时同样省略键列(见 count 源码)。
3.3 std 的数值稳定性说明
值得特别说明的是,std()在 实现文档 中明确标注采用了Welford 在线算法(单遍、数值稳定)而非 NumPy/Pandas/sklearn 常用的两遍算法,因此结果可能不同——但按官方说明,Welford 算法更准确。这与Sum/Mean等在 Arrow 引擎中通过 arrow_aggregation.py 的sum_spec、mean_spec等规格实现是同一套"累加器式(accumulator-style)"设计哲学的体现。
四、自定义聚合:aggregate() 与 AggregateFn / AggregateFnV2
内置六种聚合无法覆盖所有场景时,可以使用aggregate()传入自定义聚合器。
4.1 aggregate() 的行为与返回结构
def aggregate(self, *aggs: AggregateFn) -> Dataset: ...源码(grouped_data.py#L56-L77)表明:aggregate()接收一个或多个聚合器,输出n + 1列——第一列是分组键,第二列到第n + 1列是各聚合结果;key=None时省略键列。内部它会构造一个逻辑算子Aggregate并挂到 LogicalPlan 上,返回惰性的新Dataset,真正的计算在物化(如.take_all()、.write_parquet())时才执行。
4.2 AggregateFn:累加器式自定义聚合
aggregate.py#L70-L118 中定义的AggregateFn接受四个核心回调:
| 参数 | 作用 |
|---|---|
init | 接收分组键,返回初始累加器状态(常用 0、空 list、空 dict) |
merge | 合并两个不同 worker 产出的累加器 |
accumulate_row | 逐行更新累加器(与accumulate_block二选一) |
accumulate_block | 整块向量化更新累加器(与accumulate_row二选一) |
finalize | 可选,把最终累加器转换为输出值;缺省则原样返回 |
name | 可选,聚合结果显示列名 |
官方示例——统计每组行数的自定义聚合器:
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}]注意:AggregateFn在源码中已被标注@Deprecated(message: "AggregateFn is deprecated, please use AggregateFnV2"),新代码应优先使用AggregateFnV2。
4.3 AggregateFnV2 与内置聚合器家族
AggregateFnV2(aggregate.py#L183)是新一代聚合基类,六个内置聚合类Count、Sum、Min、Max、Mean、Std(分别定义于 aggregate.py 的第 442、590、655、731、807、894 行)均派生于它。除了这 6 个,aggregate.rst 还列出了AbsMax、Quantile、Unique、AsList、CountDistinct、ValueCounter、MissingValuePercentage、ZeroPercentage、ApproximateQuantile、ApproximateTopK等更多内置聚合,均可传入Dataset.groupby().aggregate()。
多个聚合可以一次性传入,例如:
from ray.data.aggregate import Count, Mean, Max ds.groupby("A").aggregate(Count(), Mean("B"), Max("C"))e2e 测试 test_groupby_arrow_multi_agg 与 test_groupby_multi_agg_with_nans 覆盖了这类多聚合、含 NaN 的边界场景,可作为行为参考。
五、分组变换:map_groups() 的完整参数与执行原理
map_groups()是最灵活的 GroupedData 方法:对每一个组应用一个 UDF,UDF 接收"该组全部记录构成的一个 batch",返回零个或多个记录(grouped_data.py#L98-L219)。它适合做组内归一化、组内排序、每组取 Top-K 等无法用标准聚合表达的操作。
5.1 官方示例:组内归一化
以下示例直接取自 Dataset.groupby 的文档,按variety分组后,对每组除variety外的每个特征做绝对值最大归一化:
import pandas as pd import ray def normalize_variety(group: pd.DataFrame) -> pd.DataFrame: for feature in group.drop(columns=["variety"]).columns: group[feature] = group[feature] / group[feature].abs().max() return group ds = ( ray.data.read_parquet("s3://anonymous@ray-example-data/iris.parquet") .groupby("variety") .map_groups(normalize_variety, batch_format="pandas") )另一个更简洁的示例(单记录出、单记录入):
import numpy as np ds = ray.data.from_items([ {"group": 1, "value": 1}, {"group": 1, "value": 2}, {"group": 2, "value": 3}, {"group": 2, "value": 4}, ]) # 每组返回首条记录 ds.groupby("group").map_groups( lambda g: {"result": np.array([g["value"][0]])} )5.2 完整参数说明
| 参数 | 含义与可选值 |
|---|---|
fn | 应用到每个组的函数,或可实例化的 callable 类;输入为单个组全部记录组成的 batch,输出为 0 个或多个记录(语义同map_batches()) |
zero_copy_batch | True(默认)时每组 batch 不做额外拷贝直接传入 |
compute | 计算策略。函数默认ray.data.TaskPoolStrategy()(按可用资源与输入 block 数并发起任务);TaskPoolStrategy(size=n)限制最多 n 个并发 Ray task;callable 类默认ActorPoolStrategy(min_size=1, max_size=None)自动伸缩 Actor 池,还支持size=n、min_size/max_size、initial_size等形态 |
batch_format | "default"(NumPy)、"pandas"、"pyarrow"、"cudf"(实验性)、"numpy"(Dict[str, numpy.ndarray]),或None(原样返回底层 block) |
fn_args/fn_kwargs | 传给fn的位置/关键字参数 |
fn_constructor_args/fn_constructor_kwargs | 仅当fn为 callable 类时,传给其构造函数的参数(作为 Ray actor 构造任务的顶层参数) |
num_cpus/num_gpus/memory | 每个并行 map worker 预留的 CPU / GPU / 堆内存(字节)。注意源码警告:同一任务同时指定num_cpus和num_gpus属实验特性,可能导致调度或稳定性问题 |
concurrency | 已弃用,请改用compute |
ray_remote_args_fn | 返回远程参数 dict 的函数,每次初始化 worker 前调用;返回的 args 总是覆盖ray_remote_args。已弃用,将在 Ray 2.64 移除 |
**ray_remote_args | 其他资源需求(如num_gpus=1),语义同ray.remote |
5.3 底层执行流程(源码级)
map_groups()的执行分三步(见 grouped_data.py#L222-L316):
- 数据洗牌:按前文 2.2 节的三条路径(全局
repartition(1)/ 哈希repartition+sort/ 全局sort)把相同键值的行共置; - UDF 包装:把用户函数包装成
wrapped_fn,内部调用模块级辅助函数_apply_udf_to_groups()(grouped_data.py#L614-L651)。该函数通过BlockAccessor._get_group_boundaries_sorted(keys)在已排序的 block 中定位各组边界,用slice(start, end, copy=False)切出每个组,再统一转成目标batch_format后交给 UDF;若 UDF 返回迭代器则逐个 yield,否则 yield 单结果; - map_batches_internal 执行:以
batch_size=None调用,保证每个 batch 恰好是一个完整 block(从而组不会被拆散),并以batch_format=None避免 block/batch 格式来回转换。
使用 map_groups 的两个代价(源码 grouped_data.py#L118-L125 明确提醒):
- 可能比
min()、max()等专用方法更慢; - 要求每个组能完整放进单节点内存。
因此官方建议:能用aggregate()表达的场景优先用aggregate(),map_groups()留给它真正擅长的组内变换。
六、进阶能力:with_column() 组内表达式与兼容别名
6.1 with_column():用表达式为每组新增列
with_column(column_name, expr)(grouped_data.py#L318-L381,标记为alpha稳定度)基于表达式 API 向每个组追加新列,返回保留原行列的新Dataset:
from ray.data.expressions import col ds = ( ray.data.from_items([ {"group": 1, "value": 1}, {"group": 1, "value": 2}, ]) .groupby("group") .with_column("value_twice", col("value") * 2) .sort(["group", "value"]) ) ds.take_all() # [{'group': 1, 'value': 1, 'value_twice': 2}, # {'group': 1, 'value': 2, 'value_twice': 4}]实现细节:column_name必须是非空字符串,expr必须是表达式 API 构造的Expr(否则抛TypeError),且暂不支持DownloadExpr。内部通过eval_projection([StarExpr(), aliased_expr], block)完成逐组投影,再走map_groups通道执行。
6.2 兼容别名 GroupedDataset
grouped_data.py文件末尾(第 654-655 行)保留了一个向后兼容别名:GroupedDataset = GroupedData,老代码中的GroupedDataset可直接等价使用。
七、验证与测试:从 e2e 测试看 API 边界
仓库的 test_groupby_e2e.py 提供了 100+ 个分组测试,是理解各方法边界行为的一手资料,值得关注的用例包括:
| 测试用例 | 覆盖的行为边界 |
|---|---|
test_groupby_arrow/test_groupby_tabular_count | Arrow 表格式下的基本分组与计数 |
test_groupby_none/test_groupby_map_groups_for_none_groupkey | key=None全局单组语义 |
test_groupby_errors | 非法参数(如非正num_partitions)的报错行为 |
test_groupby_tabular_sum/test_groupby_multiple_keys_tabular_count | 多列on与多列key的组合 |
test_groupby_nans/test_groupby_multi_agg_with_nans | NaN 在ignore_nulls=True/False下的不同结果 |
test_groupby_aggregations_are_associative | 聚合结果与分区方式无关(可结合性验证) |
test_groupby_map_groups_for_pandas/arrow/numpy | 不同batch_format下的 UDF 入参形态 |
test_groupby_map_groups_ray_remote_args_fn | 动态远程参数的传递 |
test_groupby_large_udf_returns | UDF 返回大数据量时的正确性 |
此外,test_infer_schema.py中的test_groupby_multi_aggs验证了多聚合结果下的 schema 推断(python/ray/data/tests/test_infer_schema.py)。
八、使用建议与注意事项
- 惰性语义:
groupby()与aggregate()都只是构建逻辑计划,不触发计算;只有消费算子(.take_all()、.to_pandas()、.write_*()等)才会真正执行,适合在 pipeline 中串联而不必担心中间结果落地。 - 优先用内置聚合:
count/sum/min/max/mean/std经由 Arrow 引擎向量化执行,性能与稳定性优于map_groups逐组调用;map_groups适合组内变换,但要保证单组内存可容纳。 - 键列合法性:
groupby(key)会对非None的 key 做 schema 校验,列不存在会直接报错;num_partitions必须为正整数。 - 空值策略:默认
ignore_nulls=True会静默忽略空值;对空值敏感的业务请显式传ignore_nulls=False以保留空值传播语义。 - API 演进:自定义聚合优先使用
AggregateFnV2(AggregateFn已弃用);concurrency与ray_remote_args_fn均已标记弃用,新代码不要使用。
如需进一步了解分组聚合的姊妹 API,可参阅 aggregate.rst(聚合器全家桶)与 dataset.rst(Dataset全部方法),以及 aggregating-data.rst 中关于数据聚合的专题讲解。
【免费下载链接】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),仅供参考