如何在Pyarrow Table中按ID统计列表类型数据的出现频率?
按ID统计PyArrow列表类型数据的出现频率问题
问题背景
需要按参与者ID统计PyArrow Table中列表类型数据的出现频率,但直接使用group_by聚合列表类型会报错。
原始数据示例
import pyarrow as pa test_table_orig = pa.table([ pa.array(["a", "a", "a", "a", "a", "b", "b", "b", "b", "b", "c", "c", "c", "c", "c", "d", "d", "d", "d", "e", "e", "e", "e", "e", "f", "f", "f", "f", "f", "f"]), pa.array([[1,1,1,1], [2,0,1,2], [3,2,1,0], [4,3,2,1], [4,3,2,1], [1,2,3,4], [1,2,3,4], [1,2,3,4], [1,2,3,4], [1,2,3,4], [5,4,3,2], [5,4,3,2], [5,4,3,2], [5,4,3,2], [4,3,2,1], [6,5,4,3], [6,5,4,3], [8,7,6,5], [9,8,7,6], [7,6,5,4], [7,6,5,4], [7,6,5,4], [7,6,5,4], [10,11,12,13], [11,12,13,14], [12,13,14,15], [33,44,55,66], [22,33,44,55], [55,66,77,88], [22,33,44,55]]) ], names=["ID", "ListData"])
直接分组报错
执行如下聚合代码:
test_table_orig.group_by(['ID','ListData']).aggregate([('ListData','count')]).to_pandas()
会抛出错误:
Traceback (most recent call last): File "<stdin>", line 1, in <module> File "pyarrow/table.pxi", line 5498, in pyarrow.lib.TableGroupBy.aggregate File "/home/usr/Python/mambaforge/lib/python3.10/site-packages/pyarrow/acero.py", line 308, in _group_by return decl.to_table(use_threads=use_threads) File "pyarrow/_acero.pyx", line 511, in pyarrow._acero.Declaration.to_table File "pyarrow/error.pxi", line 154, in pyarrow.lib.pyarrow_internal_check_status File "pyarrow/error.pxi", line 91, in pyarrow.lib.check_status pyarrow.lib.ArrowNotImplementedError: Keys of type list<item: int64>
字符串转换可行但存顾虑
将列表转为字符串后可正常统计,但实际数据中每个列表含120+元素,转换后字符串长度约2400字符,担心性能或存储问题。
解决方案
方案一:直接统计列表类型数据频率(无需转字符串)
利用PyArrow的哈希函数为列表生成唯一标量哈希值,以此作为分组依据,避免字符串转换开销:
import pyarrow.compute as pc # 为列表列生成哈希值 table_with_hash = test_table_orig.append_column("ListHash", pc.hash(test_table_orig["ListData"])) # 按ID和哈希值分组聚合,再关联回原始列表数据 result = ( table_with_hash.group_by(["ID", "ListHash"]) .aggregate([("ListData", "count")]) # 关联去重后的原始列表,恢复ListData列 .join(test_table_orig.select(["ID", "ListData", "ListHash"]).distinct(), on=["ID", "ListHash"]) .select(["ID", "ListData", "ListData_count"]) .sort_by(["ID", "ListData_count"], ascending=[True, False]) ) print(result.to_pandas())
说明:哈希值冲突概率极低,若需绝对避免,可生成双哈希值(如pc.hash(list_col, seed=1)和pc.hash(list_col, seed=2))作为联合分组键。
方案二:高效将列表转换为字符串
若倾向于字符串转换,使用PyArrow内置矢量化函数实现高效转换,避免Python循环开销:
def vectorized_list_to_str(list_col): # 用内置函数连接列表元素为字符串,再添加首尾括号 joined_elements = pc.list_join(list_col, sep=",") return pc.concat_strings(pa.array(["["]*len(list_col)), joined_elements, pa.array(["]"]*len(list_col))) # 替换列表列为字符串列 test_table_string = test_table_orig.set_column(1, "ListData", vectorized_list_to_str(test_table_orig["ListData"])) # 正常分组统计 result = test_table_string.group_by(['ID','ListData']).aggregate([('ListData','count')]) print(result.to_pandas())
说明:该方法完全基于PyArrow的C级矢量化操作,处理大数量级数据时远快于Python层面逐元素转换。
内容的提问来源于stack exchange,提问作者TDeramus
相关产品推荐
相关产品推荐

