如何将Pandas的rolling().rank(method='average')改写为Spark实现?
解决思路与代码实现
首先明确原Pandas代码的核心逻辑:按struct和otm_type分组,对指定列在过去period天的滚动窗口内计算平均排名(相同值取位置平均),再除以窗口大小得到近似百分位值。Spark实现需解决两个核心点:基于时间范围的滚动窗口、模拟Pandas的平均排名逻辑,具体步骤如下:
1. 预处理:转换时间列为数值格式
Spark的时间范围窗口依赖数值型时间表示,需先将日期列转为毫秒级时间戳(或自起始日的天数),假设你的日期列名为date_col:
from pyspark.sql import functions as F from pyspark.sql.window import Window # 将日期列转换为毫秒级时间戳 struct_ts = struct_ts.withColumn("date_ms", F.unix_timestamp("date_col") * 1000)
2. 定义基于时间的滚动窗口
使用rangeBetween而非rowsBetween,确保窗口严格对应过去period天的范围(避免因数据缺失/重复导致窗口行数不符):
# 计算period天对应的毫秒数(1天=86400秒=86400000毫秒) period_ms = period * 86400000 # 定义窗口:按分组键分区,按时间戳排序,窗口范围为过去period天到当前时间 window_spec = Window.partitionBy("struct", "otm_type")\ .orderBy("date_ms")\ .rangeBetween(-period_ms, 0)
3. 模拟Pandas的平均排名逻辑
Pandas的rank(method='average')会对窗口内相同值的位置取平均,Spark无直接对应函数,可通过两种方式实现:
方式一:精确模拟平均排名(通用场景)
对每个目标列,计算其在窗口内的最小/最大排名,取平均后除以窗口大小:
# 先计算每个窗口内的总行数 struct_ts = struct_ts.withColumn("window_size", F.count("*").over(window_spec)) # 遍历所有目标列,计算百分位值 for col in columns_to_take_pctile: # 定义按当前列排序的子窗口,用于计算排名 val_sort_window = window_spec.orderBy(col) # 计算当前值在窗口内的行号、最小排名、最大排名 struct_ts = struct_ts.withColumn(f"{col}_row_num", F.row_number().over(val_sort_window)) struct_ts = struct_ts.withColumn( f"{col}_min_rank", F.first(f"{col}_row_num").over(val_sort_window.rowsBetween(Window.unboundedPreceding, Window.currentRow)) ) struct_ts = struct_ts.withColumn( f"{col}_max_rank", F.last(f"{col}_row_num").over(val_sort_window.rowsBetween(Window.currentRow, Window.unboundedFollowing)) ) # 计算平均排名,再除以窗口大小得到最终结果 struct_ts = struct_ts.withColumn( f"{col}_percentile", (F.col(f"{col}_min_rank") + F.col(f"{col}_max_rank")) / 2 / F.col("window_size") )
方式二:简化实现(无重复值/精度要求较低场景)
如果窗口内目标列无重复值,或可接受近似结果,可利用percent_rank()转换得到近似结果:
struct_ts = struct_ts.withColumn("window_size", F.count("*").over(window_spec)) for col in columns_to_take_pctile: struct_ts = struct_ts.withColumn( f"{col}_percentile", (F.percent_rank().over(window_spec.orderBy(col)) * (F.col("window_size") - 1) + 1) / F.col("window_size") )
4. 可选:过滤窗口行数不足的记录
若需只保留窗口内至少有period天数据的记录,可添加过滤:
struct_ts = struct_ts.filter(F.col("window_size") >= period)
内容的提问来源于stack exchange,提问作者Wiktor Wijatkowski
相关产品推荐
相关产品推荐

