PySpark DataFrame计算每行最长连续1的长度问题求助
问题解决与优化方案
报错原因分析
你的方案报错主要有两个核心问题:
- UDF未指定返回类型:PySpark UDF默认返回
StringType,但你的函数返回整数,类型不匹配会引发错误。 cols变量未定义:执行F.concat(*cols)时,cols没有被赋值为目标列的列表。
另外你的findMaxConsecutiveOnes函数里,[int(j) for a,j in enumerate(X)]的enumerate是多余的,直接遍历字符串字符即可,但这不是报错原因。
修正后的UDF方案
先修复上述问题,代码如下:
from pyspark.sql import functions as F from pyspark.sql.types import IntegerType # 初始化测试数据 test_df = spark.createDataFrame([(1,1,1,0,1,1)], ("b1", "b2","b3","b4","b5","b6")) # 定义目标列列表 cols = test_df.columns # 拼接列成字符串(需将列转为字符串类型后拼接) test_df = test_df.withColumn('Ind', F.concat(*[F.col(c).cast("string") for c in cols])) # 修正后的连续1计算函数 def findMaxConsecutiveOnes(X) -> int: count = 0 maxCount = 0 for num_char in X: num = int(num_char) if num == 1: count += 1 # 每次遇到1都更新最大值,避免最后一个字符是1时遗漏 maxCount = max(maxCount, count) else: maxCount = max(maxCount, count) count = 0 return maxCount # 创建UDF并指定返回类型 maxcon_udf = F.udf(findMaxConsecutiveOnes, returnType=IntegerType()) # 计算结果 test_df = test_df.withColumn('final', maxcon_udf('Ind')) test_df.show()
执行后会输出:
+---+---+---+---+---+---+------+-----+ | b1| b2| b3| b4| b5| b6| Ind|final| +---+---+---+---+---+---+------+-----+ | 1| 1| 1| 0| 1| 1|111011| 3| +---+---+---+---+---+---+------+-----+
大数据量优化方案(无UDF)
UDF在百万级数据上性能较差,推荐使用纯PySpark内置函数实现,避免Python-JVM数据序列化开销:
from pyspark.sql import functions as F from pyspark.sql.window import Window test_df = spark.createDataFrame([(1,1,1,0,1,1)], ("b1", "b2","b3","b4","b5","b6")) cols = test_df.columns # 1. 把所有列转为数组 df_with_array = test_df.withColumn("values_array", F.array(*cols)) # 2. 展开数组为行,添加位置索引 df_exploded = df_with_array.select("*", F.posexplode("values_array").alias("pos", "val")) # 3. 标记连续1的分组:当前值为0时,分组id自增 df_grouped = df_exploded.withColumn( "group_id", F.sum(F.when(F.col("val") == 0, 1).otherwise(0)).over(Window.partitionBy(F.monotonically_increasing_id()).orderBy("pos")) ) # 4. 计算每个分组内的连续1数量,取最大值 df_result = df_grouped.groupBy(F.monotonically_increasing_id()).agg( F.max(F.when(F.col("val") == 1, F.count("*").over(Window.partitionBy(F.monotonically_increasing_id(), "group_id"))).otherwise(0)).alias("final") ) # 合并原数据和结果 final_df = test_df.join(df_result, on=test_df.monotonically_increasing_id() == df_result.monotonically_increasing_id()).drop(df_result.monotonically_increasing_id()) final_df.show()
这个方案完全基于PySpark内置函数,性能远优于UDF,适合百万级以上的数据集。
内容的提问来源于stack exchange,提问作者Lzz0
相关产品推荐
相关产品推荐

