【问题标题】:Alternative to GroupBy for Pyspark Dataframe?Pyspark Dataframe 的 GroupBy 替代方案?
【发布时间】:2020-03-04 01:11:07
【问题描述】:

我有一个这样的数据集:

timestamp     vars 
2             [1,2]
2             [1,2]
3             [1,2,3]
3             [1,2]

我想要一个这样的数据框。基本上,上述数据框中的每个值都是一个索引,该值的频率是该索引处的值。这种计算是在每个唯一的时间戳上完成的。

timestamp     vars 
2             [0, 2, 2]
3             [0,2,2,1]

现在,我按时间戳分组,并聚合/展平变量(得到类似(1,2,1,2 用于时间戳 2 或 1,2,3,1,2 用于时间戳 3)然后我有一个使用 collections.Counter 来获取 key->value dict 的 udf。然后我将这个 dict 转换为我想要的格式。

groupBy/agg 可以任意大(数组大小可以达到数百万),这似乎是 Window 函数的一个很好的用例,但我不知道如何将它们放在一起。

认为还值得一提的是,我尝试了重新分区、转换为 RDD 并使用 groupByKey。在大型数据集上,两者都非常慢(>24 小时)。

【问题讨论】:

  • 对于索引 2,它如何从 1,2,1,2 变为 [0,2,2]?带有 partitionby 子句的 windows 应该执行 groupby,如果您不使用 udf 而是使用 spark 内置函数来实现您的目标,性能可能会更好
  • 在[1,2,1,2]中有2个1和2个2。所以在索引 1 处,我输入了 2(频率),在索引 2 处,我输入了 2。由于没有 0,索引 0 仍然如此。因此,[0,2,2]。我不确定如何使用 partitionBy 和 window 从 [1,2] 和 [1,2] 到 [1,2,1,2]。试过了,但它只适用于总和。

标签: pyspark group-by pyspark-sql


【解决方案1】:

编辑: 正如 cmets 中所讨论的,原始方法的问题可能来自 count,使用了触发不必要的数据扫描的过滤器或聚合函数。下面我们在创建最终数组列之前分解数组并进行聚合(计数):

from pyspark.sql.functions import collect_list, struct  

df = spark.createDataFrame([(2,[1,2]), (2,[1,2]), (3,[1,2,3]), (3,[1,2])],['timestamp', 'vars'])

df.selectExpr("timestamp", "explode(vars) as var") \
    .groupby('timestamp','var') \
    .count() \
    .groupby("timestamp") \
    .agg(collect_list(struct("var","count")).alias("data")) \
    .selectExpr(
        "timestamp",
        "transform(data, x -> x.var) as indices",
        "transform(data, x -> x.count) as values"
    ).selectExpr(
        "timestamp",
        "transform(sequence(0, array_max(indices)), i -> IFNULL(values[array_position(indices,i)-1],0)) as new_vars"
    ).show(truncate=False)
+---------+------------+
|timestamp|new_vars    |
+---------+------------+
|3        |[0, 2, 2, 1]|
|2        |[0, 2, 2]   |
+---------+------------+

地点:

(1) 我们分解数组并对每个 timestamp + var 执行 count()

(2) groupby timestamp 并创建一个结构数组,其中包含两个字段varcount

(3) 将structs数组转换成两个数组:indices和values(类似于我们定义的SparseVector)

(4)变换序列sequence(0, array_max(indices)),对于序列中的每一个i,使用array_positionindices数组中找到i的索引,然后同时从values数组中取值位置,见下文:

IFNULL(values[array_position(indices,i)-1],0)

注意函数 array_position 使用从 1 开始的索引,而数组索引是从 0 开始的,因此我们在上面的表达式中有一个 -1

旧方法:

(1) 使用变换+滤镜/尺寸

from pyspark.sql.functions import flatten, collect_list

df.groupby('timestamp').agg(flatten(collect_list('vars')).alias('data')) \
  .selectExpr(
    "timestamp", 
    "transform(sequence(0, array_max(data)), x -> size(filter(data, y -> y = x))) as vars"
  ).show(truncate=False)
+---------+------------+
|timestamp|vars        |
+---------+------------+
|3        |[0, 2, 2, 1]|
|2        |[0, 2, 2]   |
+---------+------------+

(2)使用aggregate函数:

df.groupby('timestamp').agg(flatten(collect_list('vars')).alias('data')) \
   .selectExpr("timestamp", """ 

     aggregate(   
       data,         
       /* use an array as zero_value, size = array_max(data))+1 and all values are zero */
       array_repeat(0, int(array_max(data))+1),       
       /* increment the ith value of the Array by 1 if i == y */
       (acc, y) -> transform(acc, (x,i) -> IF(i=y, x+1, x))       
     ) as vars   

""").show(truncate=False)

【讨论】:

  • 这让我大开眼界,了解如何在转换中使用过滤器。很好的解决方案
  • 谢谢你!它有帮助,但总体上仍然会爆炸,因为该列表的长度可以达到数百万。有没有办法通过时间戳引入Window函数/分区?
  • 顺便说一句。 array_max(data) 的最大值是多少?您提到的数组中的数百万个项目是在聚合之前还是在聚合之后?
  • array_max(data) 的值可以达到一百万。变换表达式确实是大型数据集的瓶颈。需要尝试和优化它。
  • @tanyabrown,我看到了现有方法的问题,如果 M 是数组中所有项目的数量,N 是 array_max(data),那么要扫描/比较每个时间戳的数据将为 O(M*N),这对于大 M 和 N 来说效率非常低。首先分解数组进行聚合然后创建数组可能会更好。对于每一行,这可能是 O(N)。我回家后会检查这个方法。顺便提一句。是否可以为您的任务创建具有固定大小的 SparseVector 列而不是具有可变大小的 ArrayType 列?
猜你喜欢
  • 1970-01-01
  • 2018-11-14
  • 1970-01-01
  • 2022-12-22
  • 2021-08-08
  • 1970-01-01
  • 2021-09-11
  • 2019-01-01
  • 1970-01-01
相关资源
最近更新 更多