【问题标题】:Can I compute per-row aggregations over rows that satisfy a condition using PySpark?我可以使用 PySpark 计算满足条件的行的每行聚合吗?
【发布时间】:2021-02-11 22:25:43
【问题描述】:

考虑以下玩具 PySpark 数据框:

+----+-----+
|name|value|
+----+-----+
|   A|    1|
|   A|    3|
|   A|    4|
|   A|    5|
|   A|    9|
|   B|    1|
|   B|    3|
|   B|    6|
|   B|    7|
|   B|    8|
+----+-----+ 

对于每一行X,我想确定YY.name == X.name[X.value - 3, X.value + 3] 范围内有多少行Y.value。对于满足这些条件的行Y,我还想计算平均value

+----+-----+-------------+--------------+
|name|value|n_vals_in_rng|avg_val_in_rng|
+----+-----+-------------+--------------+
|   A|    1|            3|     2.6666667| # (1 + 3 + 4) / 3 = 2.6666667
|   A|    3|            4|          3.25| # (1 + 3 + 4 + 5) / 4 = 3.25
|   A|    4|            4|          3.25| # ...
|   A|    5|            3|           4.0|
|   A|    9|            1|           9.0|
|   B|    1|            2|           2.0|
|   B|    3|            3|     3.3333333|
|   B|    6|            4|           6.0|
|   B|    7|            3|           7.0|
|   B|    8|            3|           7.0|
+----+-----+-------------+--------------+

我可以在 PySpark 中有效地做到这一点吗?如果是这样,怎么做?使用 Pandas 来解决这个问题会更好吗?请注意,我的真实数据集在 name 列中有 ~400k 行和 ~8k 不同名称。


以下是我目前的解决方案。它给出了正确的结果,但在大型数据集上需要很长时间(对于具有约 400k 行的数据框需要几个小时)。

首先,我将数据框按name 分组,然后将所有values 收集到一个存储为新列的列表中。

import pyspark.sql.functions as F
import pyspark.sql.types as T
import numpy as np

# df is the data frame defined above

# define a df to be nested
df_to_nest = df.groupBy("name").agg(F.collect_list("value").alias("values"))

# df_to_nest.show():
#   +----+---------------+
#   |name|         values|
#   +----+---------------+
#   |   A|[1, 3, 4, 5, 9]|
#   |   B|[1, 3, 6, 7, 8]|
#   +----+---------------+

然后我将这个聚合数据框 (df_to_nest) 与原始 df 连接起来:

# join with the original data frame
df = df.join(df_to_nest, "name", "left")

# df.show()
#   +----+-----+---------------+
#   |name|value|         values|
#   +----+-----+---------------+
#   |   A|    1|[1, 3, 4, 5, 9]|
#   |   A|    3|[1, 3, 4, 5, 9]|
#   |   A|    4|[1, 3, 4, 5, 9]|
#   |   A|    5|[1, 3, 4, 5, 9]|
#   |   A|    9|[1, 3, 4, 5, 9]|
#   |   B|    1|[1, 3, 6, 7, 8]|
#   |   B|    3|[1, 3, 6, 7, 8]|
#   |   B|    6|[1, 3, 6, 7, 8]|
#   |   B|    7|[1, 3, 6, 7, 8]|
#   |   B|    8|[1, 3, 6, 7, 8]|
#   +----+-----+---------------+

最后,我创建一个user-defined function (UDF) 来处理每一行。

# define a UDF to process each row
def process_row(row):
    vals_in_range = [x for x in row.values if abs(x-row.value) <= 3]    
    return (len(vals_in_range),
            float(np.mean(vals_in_range)))

input_schema = F.struct([df[x] for x in df.columns])

output_schema = T.StructType([
    T.StructField("n_vals_in_rng", T.IntegerType(), nullable=True),
    T.StructField("avg_val_in_rng", T.FloatType(), nullable=True),
])

udf = F.udf(process_row, output_schema)

# apply the UDF
df= df.select("name", "value", "values", udf(input_schema).alias("new_cols"))

# unroll the new columns
df= df.select("name", "value", "new_cols.*", "values")

结果:

# df.show():
#   +----+-----+-------------+--------------+---------------+
#   |name|value|n_vals_in_rng|avg_val_in_rng|         values|
#   +----+-----+-------------+--------------+---------------+
#   |   A|    1|            3|     2.6666667|[1, 3, 4, 5, 9]|
#   |   A|    3|            4|          3.25|[1, 3, 4, 5, 9]|
#   |   A|    4|            4|          3.25|[1, 3, 4, 5, 9]|
#   |   A|    5|            3|           4.0|[1, 3, 4, 5, 9]|
#   |   A|    9|            1|           9.0|[1, 3, 4, 5, 9]|
#   |   B|    1|            2|           2.0|[1, 3, 6, 7, 8]|
#   |   B|    3|            3|     3.3333333|[1, 3, 6, 7, 8]|
#   |   B|    6|            4|           6.0|[1, 3, 6, 7, 8]|
#   |   B|    7|            3|           7.0|[1, 3, 6, 7, 8]|
#   |   B|    8|            3|           7.0|[1, 3, 6, 7, 8]|
#   +----+-----+-------------+--------------+---------------+

【问题讨论】:

  • 你在 spark 中研究过窗口函数吗?我相信这可以通过窗口函数更好地实现。

标签: python pandas dataframe pyspark apache-spark-sql


【解决方案1】:

这可以通过窗口函数和高阶函数来完成。它应该比 UDF 更有效。

df.withColumn("value", col("value").cast("int")) \
    .withColumn("values", collect_list("value").over(Window.partitionBy("name"))) \
    .withColumn("in_range", expr("filter(values, v -> abs(v - value) <= 3)")) \
    .withColumn("n_vals_in_rng", size(col("in_range"))) \
    .withColumn("avg_val_in_rng",
                expr("aggregate(in_range, 0, (acc, value) -> value + acc, acc -> acc / n_vals_in_rng)")) \
    .select("name", "value", "n_vals_in_rng", "avg_val_in_rng") \
    .show()

您可以在此处阅读有关 filteraggregate 函数的更多信息:https://spark.apache.org/docs/latest/api/sql/

【讨论】:

  • 这是一个巧妙的解决方案!请注意,它需要 Spark 2.4+
【解决方案2】:

虽然我喜欢 solution proposed by barteksch,但我发现了另一种适用于 Spark 2.3(及更低版本)的解决方案,并且更适合我的特殊情况。

我们的想法是为每对具有相同 name 且差异最多为 3 的值创建一行,而不是在列中收集范围内的值。这可以通过通过自加入和过滤。

在使用这种方法之前,请注意它可能会创建一个巨大的中间数据集,因此请确保它对您的情况有意义。对于我的大约 400k 行的数据集,爆炸后的数据集有近 4 亿行,但过滤后只保留了大约 900 万行。整个脚本运行大约需要 10 分钟。

# assign an index per row unique over a window defined by `name`
df = df.withColumn("id", F.row_number().over(Window.partitionBy("name").orderBy("value")))

# join the dataset with itself
df_exploded = df.select("name", "id", "value").join(
    df.select("name", F.col("value").alias("other_value")), "name")

# only keep rows where (value, other_value) are within 3 of each other
df_exploded = df_exploded.filter(F.abs(F.col("value") - F.col("other_value")) <= 3)

# aggregate over groups defined by (name, id)
df = df_exploded.groupBy(["name", "id"]).agg(
        # since `value` is constant per group, just take the max to retrieve it
        F.max("value").alias("value"), 
        # compute the actual aggregations: count, mean
        F.count("other_value").alias("n_vals_in_rng"),
        F.mean("other_value").alias("avg_val_in_rng")
    )

结果:

# df.show():
#   +----+---+-----+---------------+------------------+
#   |name| id|value|n_vals_in_range|  avg_val_in_range|
#   +----+---+-----+---------------+------------------+
#   |   A|  1|    1|              3|2.6666666666666665|
#   |   A|  2|    3|              4|              3.25|
#   |   A|  3|    4|              4|              3.25|
#   |   A|  4|    5|              3|               4.0|
#   |   A|  5|    9|              1|               9.0|
#   |   B|  1|    1|              2|               2.0|
#   |   B|  2|    3|              3|3.3333333333333335|
#   |   B|  3|    6|              4|               6.0|
#   |   B|  4|    7|              3|               7.0|
#   |   B|  5|    8|              3|               7.0|
#   +----+---+-----+---------------+------------------+

【讨论】:

    猜你喜欢
    • 1970-01-01
    • 1970-01-01
    • 2015-03-21
    • 1970-01-01
    • 2021-02-05
    • 2013-10-14
    • 2019-12-04
    • 1970-01-01
    • 1970-01-01
    相关资源
    最近更新 更多