如何遍历PyArrow FileSystemDataset所有分区执行聚合操作?
按分区迭代Parquet数据集并执行聚合的最优方案
原代码的核心问题
fragment.partition_expression()返回的是表达式对象(例如ds.field('year') == 2015),而非直接的分区键值,因此原代码中fragment.partition_expression() == partition_key的判断逻辑不成立,无法筛选出目标分区的片段。- FileSystemDataset 确实没有直接获取所有分区键的API,但可以通过两种可靠方式提取分区键:遍历片段解析表达式,或直接读取分区目录结构。
方案一:提取所有分区键后逐个处理(内存友好)
这种方式先获取所有唯一的年份分区,再针对每个分区单独加载数据并执行聚合,避免一次性加载全量数据集。
方法1:通过数据集片段解析分区键
import pyarrow as pa import pyarrow.dataset as ds source_path = "你的数据集路径" # 初始化分区化数据集 dataset = ds.dataset( source_path, format="parquet", partitioning=ds.partitioning(pa.schema([("year", pa.int16())])) ) # 提取所有唯一的年份分区值 unique_years = set() for fragment in dataset.get_fragments(): expr = fragment.partition_expression() # 跳过无分区的片段(如果存在) if not expr or expr.type == pa.scalar(True).type: continue # 解析表达式中的年份值 year_val = expr.right.value.as_py() unique_years.add(year_val) # 遍历每个年份分区执行聚合 for year in sorted(unique_years): # 用Scanner过滤当前年份的所有数据,支持谓词下推 scanner = dataset.scanner( filter=ds.field("year") == year, columns=["agg_col", "value_col"] # 只加载需要的列,进一步节省内存 ) # 加载当前分区数据并执行聚合 partition_table = scanner.to_table() agg_result = partition_table.group_by("agg_col").sum("value_col") # 处理聚合结果(示例:转为Pandas DataFrame查看) print(f"===== 年份 {year} 聚合结果 =====") print(agg_result.to_pandas())
方法2:直接遍历分区目录(更高效)
如果你的Parquet数据集使用标准的Hive风格分区目录(例如 year=2015/、year=2016/),可以直接通过文件系统遍历目录提取分区键,效率比遍历片段更高:
import pyarrow as pa import pyarrow.dataset as ds from pyarrow import fs source_path = "你的数据集路径" fs_handle = fs.LocalFileSystem() # 若使用S3,替换为fs.S3FileSystem() # 遍历根目录下的分区目录 unique_years = set() for item in fs_handle.get_file_info(fs.FileSelector(source_path, recursive=False)): if item.is_dir and item.base_name.startswith("year="): # 从目录名中解析年份值 year = int(item.base_name.split("=")[1]) unique_years.add(year) # 后续聚合逻辑同方法1 dataset = ds.dataset( source_path, format="parquet", partitioning=ds.partitioning(pa.schema([("year", pa.int16())])) ) for year in sorted(unique_years): scanner = dataset.scanner(filter=ds.field("year") == year, columns=["agg_col", "value_col"]) partition_table = scanner.to_table() agg_result = partition_table.group_by("agg_col").sum("value_col") print(f"===== 年份 {year} 聚合结果 =====") print(agg_result.to_pandas())
方案二:直接按分区批量聚合(最简洁)
如果不需要逐个分区单独处理,而是希望直接得到按年份+聚合列的汇总结果,可以使用PyArrow Dataset的aggregate方法,底层会自动按分区处理,无需手动遍历:
import pyarrow as pa import pyarrow.dataset as ds source_path = "你的数据集路径" dataset = ds.dataset( source_path, format="parquet", partitioning=ds.partitioning(pa.schema([("year", pa.int16())])) ) # 定义聚合规则:按year和agg_col分组,对value_col求和 aggregations = [ds.sum("value_col").alias("total_value")] final_result = dataset.aggregate( group_by=["year", "agg_col"], aggregations=aggregations ) # 输出最终聚合结果 print(final_result.to_pandas())
关键优化点
- 使用
scanner的columns参数指定仅加载需要的列,避免加载无关数据占用内存。 - 优先使用文件系统遍历目录的方式提取分区键,尤其是数据集规模较大时,比遍历片段更快。
- 利用PyArrow的内置聚合API(
group_by/aggregate),比手动拼接表后聚合更高效。
内容的提问来源于stack exchange,提问作者Stephen Nicholson
相关产品推荐
相关产品推荐

