如何计算Spark DataFrame数组列中连续相同整数的最大个数?
问题:计算Spark DataFrame数组列中连续相同元素的最大次数
原始DataFrame
df = spark.createDataFrame([ [0, [1, 1, 4, 4, 4]], [1, [3, 2, 2, -4]], [2, [1, 1, 5, 5]], [3, [-1, -9, -9, -9, -9]]] , ['id', 'array_col'] ) df.show()
输出:
+---+--------------------+ | id| array_col| +---+--------------------+ | 0| [1, 1, 4, 4, 4]| | 1| [3, 2, 2, -4]| | 2| [1, 1, 5, 5]| | 3|[-1, -9, -9, -9, -9]| +---+--------------------+
期望结果
+---+--------------------+-------------------------+ | id| array_col|max_consecutive_identical| +---+--------------------+-------------------------+ | 0| [1, 1, 4, 4, 4]| 3| | 1| [3, 2, 2, -4]| 2| | 2| [1, 1, 5, 5]| 2| | 3|[-1, -9, -9, -9, -9]| 4| +---+--------------------+-------------------------+
尝试的方法及问题
尝试将数组拼接为字符串,用regexp_extract_all处理,但遇到多个问题:
from pyspark.sql.functions import col, concat_ws, expr df = df.withColumn('joined_str', concat_ws('', col('array_col'))) df.show()
输出:
+---+--------------------+-----------+ | id| array_col| joined_str| +---+--------------------+-----------+ | 0| [1, 1, 4, 4, 4]| 11444| | 1| [3, 2, 2, -4]| 322-4| | 2| [1, 1, 5, 5]| 1155| | 3|[-1, -9, -9, -9, -9]| -1-9-9-9-9| +---+--------------------+-----------+
df = df.withColumn('regexp_extracted', expr('regexp_extract_all(joined_str, "([0-9])\\1*", 1)')) df.show()
输出:
+---+--------------------+----------+----------+----------------+ | id| array_col| concat_ws|joined_str|regexp_extracted| +---+--------------------+----------+----------+----------------+ | 0| [1, 1, 4, 4, 4]| 11444| 11444| [1, 1, 4, 4, 4]| | 1| [3, 2, 2, -4]| 322-4| 322-4| [3, 2, 2, 4]| | 2| [1, 1, 5, 5]| 1155| 1155| [1, 1, 5, 5]| | 3|[-1, -9, -9, -9, -9]|-1-9-9-9-9|-1-9-9-9-9| [1, 9, 9, 9, 9]| +---+--------------------+----------+----------+----------------+
遇到的问题:
- 负数匹配错误:
-4被提取为4,-1被提取为1,丢失负号 - 多位数无法处理:若数组包含多位数,拼接后会与其他数字混淆,无法识别完整元素
- 正则逻辑失效:即使是个位数,正则也没合并连续相同元素,反而拆成了单个元素
解决方案
放弃字符串拼接+正则的思路,改用Spark窗口函数和分组统计,精准处理连续元素:
from pyspark.sql import Window from pyspark.sql.functions import col, posexplode, lag, when, count, max # 1. 展开数组,获取每个元素的位置索引 exploded_df = df.select( col('id'), col('array_col'), posexplode(col('array_col')).alias('pos', 'value') ) # 2. 生成分组标识:当前元素与前一个不同时,分组ID递增 window_spec = Window.partitionBy('id').orderBy('pos') grouped_df = exploded_df.withColumn( 'group_id', count(when(col('value') != lag(col('value')).over(window_spec), 1)).over(window_spec) ) # 3. 统计每个分组的连续次数,再取每个ID的最大值 result_df = grouped_df.groupBy('id', 'array_col', 'group_id')\ .agg(count('*').alias('consecutive_count'))\ .groupBy('id', 'array_col')\ .agg(max('consecutive_count').alias('max_consecutive_identical')) result_df.show()
输出:
+---+--------------------+-------------------------+ | id| array_col|max_consecutive_identical| +---+--------------------+-------------------------+ | 0| [1, 1, 4, 4, 4]| 3| | 1| [3, 2, 2, -4]| 2| | 2| [1, 1, 5, 5]| 2| | 3|[-1, -9, -9, -9, -9]| 4| +---+--------------------+-------------------------+
方案说明
- 彻底规避字符串拼接的局限性,能正确处理负数、多位数等任意合法数值元素
- 通过窗口函数精准识别连续相同元素的分组,再统计每组长度,逻辑清晰可靠
内容的提问来源于stack exchange,提问作者L.B.
相关产品推荐
相关产品推荐

