You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.05.14 06:45:10