将T-SQL SEGMENTATION存储过程转换为PySpark函数遇问题求助
原T-SQL存储过程转PySpark函数实现及优化
原T-SQL核心逻辑梳理
- 基于指定列排序计算累计和
CUM_SUM - 可选异常值处理:若
OUTLIER不等于1000000,扣除所有超过阈值的目标列总和,调整累计和 - 根据累计和占最大值的比例计算分段值
COL_NM2,规则为:目标列>0时11-CEILING(10*CUM_SUM/@MAX),否则为0 - 将分段值关联回原表,关联键为
PCYC_ID和NPI(NULL替换为0) - 修正分段值:将等于11的数值改为10
PySpark实现代码
from pyspark.sql import functions as F from pyspark.sql.window import Window def segmentation(col_nm, input_df, col_nm2, row1, row2, row3, row4, row5, row6, outlier): # 1. 计算排序后的累计和CUM_SUM window_spec = Window.orderBy( F.col(row1).desc(), F.col(row2).desc(), F.col(row3).desc(), F.col(row4).desc(), F.col(row5).desc(), F.col(row6).desc() ).rowsBetween(Window.unboundedPreceding, Window.currentRow) tmp_df = input_df.withColumn("CUM_SUM", F.sum(col_nm).over(window_spec)) # 2. 异常值处理逻辑 if str(outlier) != "1000000": outlier_sum_row = input_df.filter(F.col(col_nm) >= F.lit(outlier)).agg(F.sum(col_nm).alias("sum_outlier")).first() outlier_sum = outlier_sum_row["sum_outlier"] if outlier_sum_row else 0 if outlier_sum > 0: tmp_df = tmp_df.withColumn("CUM_SUM", F.col("CUM_SUM") - F.lit(outlier_sum)) # 3. 计算分段值COL_NM2 max_cum_row = tmp_df.agg(F.max("CUM_SUM").alias("max_cum")).first() max_cum_sum = max_cum_row["max_cum"] if max_cum_row else 0 if max_cum_sum == 0: tmp_df = tmp_df.withColumn(col_nm2, F.when(F.col(col_nm) > 0, F.lit(0)).otherwise(F.lit(0))) else: tmp_df = tmp_df.withColumn( col_nm2, F.when( F.col(col_nm) > 0, 11 - F.ceil((10 * F.col("CUM_SUM")) / F.lit(max_cum_sum)) ).otherwise(F.lit(0)) ) # 4. 关联原表,更新分段值 join_keys = [ F.coalesce(input_df["PCYC_ID"], F.lit(0)) == F.coalesce(tmp_df["PCYC_ID"], F.lit(0)), F.coalesce(input_df["NPI"], F.lit(0)) == F.coalesce(tmp_df["NPI"], F.lit(0)) ] result_df = input_df.join(tmp_df.select("PCYC_ID", "NPI", col_nm2), on=join_keys, how="inner") # 5. 修正分段值:11改为10 result_df = result_df.withColumn( col_nm2, F.when(F.col(col_nm2) == 11, F.lit(10)).otherwise(F.col(col_nm2)) ) return result_df
使用示例
# 读取输入表(以Hive表为例) input_df = spark.read.table("your_input_table_name") # 调用分段函数 result_df = segmentation( col_nm="target_column", input_df=input_df, col_nm2="segment_column", row1="sort_col1", row2="sort_col2", row3="sort_col3", row4="sort_col4", row5="sort_col5", row6="sort_col6", outlier="500000" ) # 将结果写回原表(覆盖模式) result_df.write.mode("overwrite").saveAsTable("your_input_table_name")
性能优化要点
- 减少重复扫描:复用输入DataFrame,避免原T-SQL中多次读取表的开销
- 窗口函数优化:若数据量极大,提前按
row1-row6分区,缩小窗口函数的Shuffle范围 - 广播小数据集:异常值计算结果若很小,使用
F.broadcast()广播outlier_sum,减少关联Shuffle - 空值与边界防护:补充了
max_cum_sum为0的处理逻辑,避免运行时除以0错误
内容的提问来源于stack exchange,提问作者DEVEN MALI
相关产品推荐
相关产品推荐

