Spark百万行DataFrame:按ID统计日期在起止区间内的行数
解决方案:为Spark DataFrame添加sum_of_rows列
针对你需要处理百万行Spark DataFrame的需求,这里提供一个高效的解决方案,避免了性能低下的自连接操作,同时完全匹配你的业务规则:
步骤1:数据准备与类型转换
首先得把原始的字符串日期转换成Spark的DateType,这样才能正确进行日期大小比较:
from pyspark.sql import SparkSession from pyspark.sql import functions as F from pyspark.sql.types import DateType # 初始化SparkSession(如果还没创建) spark = SparkSession.builder.appName("SumOfRowsCalculation").getOrCreate() # 创建示例表 table = spark.createDataFrame( [ ["A", '2008-01-02', '2010-01-01', '2009-01-01'], ["A", '2005-01-02', '2012-01-01', None], ["A", '2013-01-02', '2015-01-01', '2014-01-01'], ["B", '2002-01-02', '2019-01-01', '2003-01-01'], ["B", '2015-01-02', '2017-01-01', '2016-01-01'] ], ("Id", "start", "end", "date") ) # 转换字符串字段为日期类型 table = table.withColumn("start", F.col("start").cast(DateType())) \ .withColumn("end", F.col("end").cast(DateType())) \ .withColumn("date", F.col("date").cast(DateType()))
步骤2:高效计算sum_of_rows
核心思路是按Id分组收集所有时间区间,再对每行的date过滤匹配的区间并计数,这种方法的时间复杂度远低于自连接,非常适合百万级数据:
方案1:Spark 3.x+ 内置函数(推荐)
Spark 3.x及以上支持F.filter内置函数,无需自定义UDF,性能更优:
# 按Id分组,收集该Id下所有的时间区间为结构体数组 grouped_intervals = table.groupBy("Id").agg( F.collect_list(F.struct("start", "end")).alias("interval_list") ) # 关联原表并计算sum_of_rows result_df = table.join(grouped_intervals, on="Id", how="left") \ .withColumn( "sum_of_rows", F.when( F.col("date").isNotNull(), # 过滤出满足start ≤ date ≤ end的区间,再统计数量 F.size(F.filter(F.col("interval_list"), lambda x: x.start <= F.col("date") <= x.end)) ).otherwise(None) # date为null时返回null ) \ .drop("interval_list") # 移除临时列 # 查看结果 result_df.show()
方案2:Spark 2.x 兼容(自定义UDF)
如果你的Spark版本低于3.x,可以用自定义UDF实现相同逻辑:
from pyspark.sql.functions import udf from pyspark.sql.types import IntegerType def count_matching_intervals(intervals, target_date): if target_date is None: return None match_count = 0 for interval in intervals: if interval["start"] <= target_date <= interval["end"]: match_count += 1 return match_count # 注册UDF count_intervals_udf = udf(count_matching_intervals, IntegerType()) # 关联并计算 result_df = table.join(grouped_intervals, on="Id", how="left") \ .withColumn( "sum_of_rows", count_intervals_udf(F.col("interval_list"), F.col("date")) ) \ .drop("interval_list") result_df.show()
结果验证
运行上述代码后,你会得到完全符合预期的输出:
+---+----------+----------+----------+-----------+ | Id| start| end| date|sum_of_rows| +---+----------+----------+----------+-----------+ | A|2008-01-02|2010-01-01|2009-01-01| 2| | A|2005-01-02|2012-01-01| null| null| | A|2013-01-02|2015-01-01|2014-01-01| 1| | B|2002-01-02|2019-01-01|2003-01-01| 1| | B|2015-01-02|2017-01-01|2016-01-01| 2| +---+----------+----------+----------+-----------+
性能说明
- 该方案避免了自连接带来的O(N*K)时间复杂度(N为总行数,K为每个Id的平均行数),而是采用O(N)的分组收集 + 每行O(K)的过滤,整体性能更适合百万级数据;
- 内置函数方案比UDF方案性能更高,因为Spark可以对内置函数进行深度优化,而UDF需要跨语言数据传输(如果使用PySpark的话)。
内容的提问来源于stack exchange,提问作者Peter Martins
相关产品推荐
相关产品推荐

