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
相关产品推荐
相关产品推荐

