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

PySpark/Databricks中高效并行计算时序Kolmogorov-Smirnov检验

问题描述

拥有包含snapshot(时间列)和values列的PySpark DataFrame,需对每个snapshot的数值与前一个snapshot的数值执行Kolmogorov-Smirnov(KS)检验。初始用for循环实现,但数据量极大时效率极低;尝试UDF实现,数据集规模增大后出现序列化错误,寻求可利用集群Worker节点、避免循环的高效并行解决方案。

数据示例

+----------+--------------------+
|  snapshot|              values|
+----------+--------------------+
|2005-01-31| 0.19120256617637743|
|2005-01-31|  0.7972692479278891|
|2005-02-28|0.005236883665445502|
|2005-02-28|  0.5474099672222935|
|2005-02-28| 0.13077227571485905|
+----------+--------------------+

初始for循环实现

import numpy as np
from scipy.stats import ks_2samp
import pyspark.sql.functions as F

def KS_for_one_snapshot(temp_df, snapshots_list, j, var = "values"):
  sample1=temp_df.filter(F.col("snapshot")==snapshots_list[j])
  sample2=temp_df.filter(F.col("snapshot")==snapshots_list[j-1]) # 与前一个snapshot对比

  if (sample1.count() == 0 or sample2.count() == 0 ):
    ks_value = -1  # 避免空样本导致类型错误
  else:
    ks_value, p_value = ks_2samp(  np.array(sample1.select(var).collect()).reshape(-1)
                                 , np.array(sample2.select(var).collect()).reshape(-1)
                                 , alternative="two-sided"
                                 , mode="auto")


  return ks_value

results = []
snapshots_list = df.select('snapshot').dropDuplicates().sort('snapshot').rdd.flatMap(lambda x: x).collect()
for j in range(len(snapshots_list) - 1 ):
  results.append(KS_for_one_snapshot(df, snapshots_list, j+1))
results

数据生成代码

import pyspark.sql.types as T
from random import randint

df = (spark.createDataFrame( range(1,1000), T.IntegerType())
      .withColumn('snapshot' ,F.array(F.lit("2005-01-31"), F.lit("2005-02-28"),F.lit("2005-03-30")  ).getItem((F.rand()*3).cast("int")))
      .withColumn('values', F.rand()).drop('value')
)

尝试的UDF实现及报错

UDF代码

var_used = 'values'
data_input_1 = df.groupBy('snapshot').agg(F.collect_list(var_used).alias('value_list'))
data_input_2 = df.groupBy('snapshot').agg(F.collect_list(var_used).alias("value_list_2"))
windowSpec = Window.orderBy("snapshot")
data_input_2 = data_input_2.withColumn('snapshot_2', F.lag("snapshot", 1).over(Window.orderBy('snapshot'))).filter('snapshot_2 is not NULL')
data_input_final = data_input_1.join(data_input_2, data_input_1.snapshot == data_input_2.snapshot_2)

def KS_one_snapshot_general(sample_in_list_1, sample_in_list_2):
  
  if (len(sample_in_list_1) == 0 or len(sample_in_list_2) == 0 ):
    ks_value = -1  # 避免空样本导致类型错误
  else:
    ks_value, p_value = ks_2samp(  sample_in_list_1
                                 , sample_in_list_2
                                 , alternative="two-sided"
                                 , mode="auto")
  
  return ks_value

import pyspark.sql.types as T
KS_one_snapshot_general_udf = udf(KS_one_snapshot_general, T.FloatType())
data_input_final.select( KS_one_snapshot_general_udf('value_list', 'value_list_2')).display()

报错信息(翻译后)

PickleException: 构造ClassDict时预期零参数(针对numpy.dtype)

高效并行解决方案

普通UDF报错是因为scipy返回的numpy类型无法被Spark默认的Pickle序列化器正确处理,且循环方案是Driver端串行执行,完全未利用集群资源。推荐使用Pandas UDF(Scalar UDF),它基于Arrow序列化,避免Pickle问题,且能在Worker节点并行执行。

完整实现代码

import pyspark.sql.functions as F
import pyspark.sql.types as T
import pandas as pd
from scipy.stats import ks_2samp
from pyspark.sql.functions import pandas_udf

# 1. 按snapshot分组收集values列表
grouped_df = df.groupBy("snapshot").agg(F.collect_list("values").alias("current_values"))

# 2. 用窗口函数获取前一个snapshot的values列表
window_spec = Window.orderBy("snapshot")
paired_df = grouped_df.withColumn(
    "prev_values",
    F.lag("current_values", 1).over(window_spec)
).filter(F.col("prev_values").isNotNull())  # 过滤第一个无前置的snapshot

# 3. 定义Pandas Scalar UDF处理KS检验
@pandas_udf(T.FloatType())
def calculate_ks(current_vals: pd.Series, prev_vals: pd.Series) -> pd.Series:
    def ks_pair(curr, prev):
        if len(curr) == 0 or len(prev) == 0:
            return -1.0
        ks_stat, _ = ks_2samp(curr, prev, alternative="two-sided", mode="auto")
        return float(ks_stat)
    
    # 对每一对列表执行KS检验
    return pd.Series([ks_pair(c, p) for c, p in zip(current_vals, prev_vals)])

# 4. 执行计算并得到结果
result_df = paired_df.select(
    "snapshot",
    F.lag("snapshot", 1).over(window_spec).alias("prev_snapshot"),
    calculate_ks("current_values", "prev_values").alias("ks_statistic")
)

result_df.display()

方案说明

  • 分组收集:通过groupBy+collect_list将每个snapshot的values聚合为列表,减少数据 shuffle 次数
  • 窗口关联:用lag窗口函数直接关联当前与前一个snapshot的values列表,避免低效的join操作
  • Pandas UDF:基于Arrow序列化,解决numpy类型的Pickle问题,同时利用Spark的并行计算能力,将KS检验任务分发到各个Worker节点执行
  • 空值处理:保留了空样本返回-1的逻辑,避免运行时错误

内容的提问来源于stack exchange,提问作者George Sotiropoulos

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.23 00:18:19