如何在PySpark中按列聚合数组元素?
输入DataFrame结构如下:
| col1 | col2 |
|---|---|
| [1,2,3,4] | [0,1,0,3] |
| [5,6,7,8] | [0,3,4,8] |
期望输出结果:
| col1 | col2 |
|---|---|
| [6,8,10,12] | [0,4,4,11] |
在Snowflake Snowpark中,用array_construct可以直接实现这类数组对应位置聚合的需求,比如类似写法:array_construct(count('*'), sum(col('x')), sum(col('y')), count(col('y')))。但在Apache Spark中,对应的array()函数会被识别为聚合函数,直接嵌套聚合操作会抛出错误:
pyspark.sql.utils.AnalysisException: It is not allowed to use an aggregate function in the argument of another aggregate function. Please use the inner aggregate function in a sub-query.;
现在需要编写同时兼容Snowpark和Spark的代码,目前考虑过groupby + collect_list的思路,下面是几种可行的解决方案:
方案1:拆列聚合后重组数组
这是兼容性最强的方案,核心思路是先把数组的每个位置拆成独立列,聚合后再重新组合成数组。
实现步骤
- 确认所有行的数组长度一致(若长度不固定,可先通过
size函数动态获取) - 对目标列的数组按位置提取元素,生成临时列
- 对临时列执行聚合操作(比如求和)
- 将聚合后的临时列重新组合成数组列
Spark示例代码
from pyspark.sql import functions as F # 假设数组长度为4,动态获取可替换为df.select(F.size(F.col("col1"))).collect()[0][0] array_length = 4 # 拆分数组元素为临时列 df = df.withColumn("col1_0", F.element_at(F.col("col1"), 1)) \ .withColumn("col1_1", F.element_at(F.col("col1"), 2)) \ .withColumn("col1_2", F.element_at(F.col("col1"), 3)) \ .withColumn("col1_3", F.element_at(F.col("col1"), 4)) \ .withColumn("col2_0", F.element_at(F.col("col2"), 1)) \ .withColumn("col2_1", F.element_at(F.col("col2"), 2)) \ .withColumn("col2_2", F.element_at(F.col("col2"), 3)) \ .withColumn("col2_3", F.element_at(F.col("col2"), 4)) # 聚合求和 agg_df = df.agg( F.sum("col1_0").alias("sum_col1_0"), F.sum("col1_1").alias("sum_col1_1"), F.sum("col1_2").alias("sum_col1_2"), F.sum("col1_3").alias("sum_col1_3"), F.sum("col2_0").alias("sum_col2_0"), F.sum("col2_1").alias("sum_col2_1"), F.sum("col2_2").alias("sum_col2_2"), F.sum("col2_3").alias("sum_col2_3") ) # 重组数组并清理临时列 result_df = agg_df.withColumn("col1", F.array("sum_col1_0", "sum_col1_1", "sum_col1_2", "sum_col1_3")) \ .withColumn("col2", F.array("sum_col2_0", "sum_col2_1", "sum_col2_2", "sum_col2_3")) \ .drop(*[c for c in agg_df.columns if c.startswith("sum_")])
Snowpark示例代码
逻辑与Spark完全一致,仅需将array替换为array_construct:
from snowflake.snowpark import functions as F array_length = 4 df = df.withColumn("col1_0", F.element_at(F.col("col1"), 1)) \ .withColumn("col1_1", F.element_at(F.col("col1"), 2)) \ .withColumn("col1_2", F.element_at(F.col("col1"), 3)) \ .withColumn("col1_3", F.element_at(F.col("col1"), 4)) \ .withColumn("col2_0", F.element_at(F.col("col2"), 1)) \ .withColumn("col2_1", F.element_at(F.col("col2"), 2)) \ .withColumn("col2_2", F.element_at(F.col("col2"), 3)) \ .withColumn("col2_3", F.element_at(F.col("col2"), 4)) agg_df = df.agg( F.sum("col1_0").alias("sum_col1_0"), F.sum("col1_1").alias("sum_col1_1"), F.sum("col1_2").alias("sum_col1_2"), F.sum("col1_3").alias("sum_col1_3"), F.sum("col2_0").alias("sum_col2_0"), F.sum("col2_1").alias("sum_col2_1"), F.sum("col2_2").alias("sum_col2_2"), F.sum("col2_3").alias("sum_col2_3") ) result_df = agg_df.withColumn("col1", F.array_construct("sum_col1_0", "sum_col1_1", "sum_col1_2", "sum_col1_3")) \ .withColumn("col2", F.array_construct("sum_col2_0", "sum_col2_1", "sum_col2_2", "sum_col2_3")) \ .drop(*[c for c in agg_df.columns if c.startswith("sum_")])
方案2:收集数组列表后逐元素聚合
适合数组长度不固定的场景,先收集所有行的数组到列表,再用高阶函数对列表中的数组逐元素执行聚合。
Spark示例代码
from pyspark.sql import functions as F # 收集所有行的数组到列表 collected_df = df.agg( F.collect_list("col1").alias("col1_list"), F.collect_list("col2").alias("col2_list") ) # 用transform+aggregate实现逐元素求和 result_df = collected_df.withColumn( "col1", F.transform( F.sequence(F.lit(0), F.size(F.col("col1_list")[0])-1), lambda i: F.aggregate( F.col("col1_list"), F.lit(0), lambda acc, arr: acc + arr[i] ) ) ).withColumn( "col2", F.transform( F.sequence(F.lit(0), F.size(F.col("col2_list")[0])-1), lambda i: F.aggregate( F.col("col2_list"), F.lit(0), lambda acc, arr: acc + arr[i] ) ) ).drop("col1_list", "col2_list")
Snowpark示例代码
from snowflake.snowpark import functions as F # 收集所有行的数组到列表 collected_df = df.agg( F.array_agg("col1").alias("col1_list"), F.array_agg("col2").alias("col2_list") ) # 动态获取数组长度 array_len = collected_df.select(F.size(F.col("col1_list")[0])).collect()[0][0] # 逐元素求和并重组数组 result_df = collected_df.withColumn( "col1", F.array_construct(*[ F.aggregate(F.col("col1_list"), F.lit(0), lambda acc, arr: acc + F.element_at(arr, i+1)) for i in range(array_len) ]) ).withColumn( "col2", F.array_construct(*[ F.aggregate(F.col("col2_list"), F.lit(0), lambda acc, arr: acc + F.element_at(arr, i+1)) for i in range(array_len) ]) ).drop("col1_list", "col2_list")
方案3:封装通用兼容函数
可以将两种引擎的实现封装成统一函数,通过判断DataFrame类型自动选择逻辑,上层业务代码无需关心底层引擎:
def aggregate_array(df, col_name, agg_func): # 判断当前环境:Spark DataFrame带有_jdf属性 if hasattr(df, '_jdf'): # Spark逻辑:动态拆列聚合 array_len = df.select(F.size(F.col(col_name))).collect()[0][0] # 拆分临时列 for i in range(array_len): df = df.withColumn(f"{col_name}_{i}", F.element_at(F.col(col_name), i+1)) # 聚合 agg_cols = [agg_func(f"{col_name}_{i}").alias(f"agg_{col_name}_{i}") for i in range(array_len)] agg_df = df.agg(*agg_cols) # 重组数组 return agg_df.withColumn(col_name, F.array(*[f"agg_{col_name}_{i}" for i in range(array_len)])).drop(*agg_cols) else: # Snowpark逻辑:直接用array_construct嵌套聚合 array_len = df.select(F.size(F.col(col_name))).collect()[0][0] return df.agg( F.array_construct(*[ agg_func(F.element_at(F.col(col_name), i+1)) for i in range(array_len) ]).alias(col_name) ) # 使用示例:对col1和col2分别求和 result_df = aggregate_array(df, "col1", F.sum).join(aggregate_array(df, "col2", F.sum), how="cross")
内容的提问来源于stack exchange,提问作者whisperstream

