Databricks上PySpark代码执行停滞,求SKU时间窗口计数器优化方案
问题描述
我有两个DataFrame:df_selected 和 df_filtered_mins_60,结构与数据量如下:
df_filtered_mins_60.columns输出:["CSku", "start_timestamp", "end_timestamp"],共112,397行df_selected.columns输出:["DATEUPDATED", "DATE", "HOUR", "CPSKU", "BB_Status", "ActivePrice", "PrevPrice", "MinPrice", "AsCost", "MinMargin", "CPT", "Comp_Price", "AP_MSG"],共7,816,521行
需求
遍历df_filtered_mins_60的每一行,获取sku=CSku、start_time=start_timestamp、stop_time=end_timestamp;为df_selected中满足DATEUPDATED在[start_time, stop_time]区间且CPSKU等于sku的行分配常量i,每处理完一行df_filtered_mins_60,i自增1。后续将通过该计数器筛选各SKU对应时间窗口内的行,进行聚合分析。
原代码问题
我编写的代码运行时会卡住数小时,直至手动终止:
i = 1 df_selected = df_selected.withColumn("counter", lit(0)) # 遍历df_filtered_mins_60的每一行 for row in df_filtered_mins_60.collect(): sku = row['CSku'] start_time = row['start_timestamp'] stop_time = row['stop_timestamp'] # 筛选条件并更新counter列 df_selected = df_selected.withColumn("counter", when((df_selected.DATEUPDATED >= start_time) & (df_selected.DATEUPDATED <= stop_time) & (df_selected.CPSKU == sku), lit(i)).otherwise(df_selected.counter)) i += 1 # 展示更新后的DataFrame display(df_selected)
示例数据生成代码
df_selected生成代码
from pyspark.sql import SparkSession from pyspark.sql.functions import to_timestamp from pyspark.sql.types import StructType, StructField, StringType, IntegerType, DoubleType, TimestampType # 初始化SparkSession spark_a = SparkSession.builder \ .appName("Create DataFrame") \ .getOrCreate() schema = StructType([ StructField("DATEUPDATED", StringType(), True), StructField("DATE", StringType(), True), StructField("HOUR", IntegerType(), True), StructField("CPSKU", StringType(), True), StructField("BB_Status", IntegerType(), True), StructField("ActivePrice", DoubleType(), True), StructField("PrevPrice", DoubleType(), True), StructField("MinPrice", DoubleType(), True), StructField("AsCost", DoubleType(), True), StructField("MinMargin", DoubleType(), True), StructField("CPT", DoubleType(), True), StructField("Comp_Price", DoubleType(), True) ]) data=[('2024-01-01T19:45:39.151+00:00','2024-01-01',0,'MSAN10115836',0,14.86,14.86,14.86,12.63,0.00,13.90,5.84) , ('2024-01-01T19:55:10.904+00:00','2024-01-01',0,'MSAN10115836',0,126.04,126.04,126.04,108.96,0.00,0.00,93.54), ('2024-01-01T20:35:10.904+00:00','2024-01-01',0,'MSAN10115836',0,126.04,126.04,126.04,108.96,0.00,0.00,93.54), ('2024-01-15T12:55:18.528+00:00','2024-01-01',1,'PFXNDDF4OX',1,18.16,18.16,10.56,26.85,-199.00,18.16,34.10) , ('2024-01-15T13:25:18.528+00:00','2024-01-01',1,'PFXNDDF4OX',1,18.16,18.16,10.56,26.85,-199.00,18.16,34.10) , ('2024-01-15T13:35:18.528+00:00','2024-01-01',1,'PFXNDDF4OX',1,18.16,18.16,10.56,26.85,-199.00,18.16,34.10) , ('2024-01-15T13:51:09.574+00:00','2024-01-01',1,'PFXNDDF4OX',1,20.16,18.16,10.56,26.85,-199.00,18.16,34.10) , ('2024-01-15T07:28:48.265+00:00','2024-01-01',1,'DEWNDCB135C',0,44.93,44.93,44.93,38.09,0.25,26.9,941.26), ('2024-01-15T07:50:32.412+00:00','2024-01-01',1,'DEWNDCB135C',0,44.93,44.93,44.93,38.09,0.25,26.9,941.26), ('2024-01-15T07:52:32.412+00:00','2024-01-01',1,'DEWNDCB135C',0,44.93,44.93,44.93,38.09,0.25,26.9,941.26)] df_selected = spark.createDataFrame(data, schema=schema) df_selected = df_selected.withColumn("DateUpdated", to_timestamp(df_selected["DATEUPDATED"], "yyyy-MM-dd'T'HH:mm:ss.SSS'+00:00'")) display(df_selected)
df_filtered_mins_60生成代码
schema = StructType([ StructField("CPSKU", StringType(), True), StructField("start_timestamp", StringType(), True), StructField("stop_timestamp", StringType(), True) ]) data_2=[('MSAN10115836','2024-01-01T19:45:39.151+00:00','2024-01-01T20:35:10.904+00:00'), ('MSAN10115836','2024-01-08T06:04:16.484+00:00','2024-01-08T06:42:14.912+00:00'), ('DEWNDCB135C','2024-01-15T07:28:48.265+00:00','2024-01-15T07:52:32.412+00:00'), ('DEWNDCB135C','2024-01-15T11:37:56.698+00:00','2024-01-15T12:35:09.693+00:00'), ('PFXNDDF4OX','2024-01-15T12:55:18.528+00:00','2024-01-15T13:51:09.574+00:00'), ('PFXNDDF4OX','2024-01-15T19:25:10.150+00:00','2024-01-15T20:24:36.385+00:00')] df_filtered_mins_60 = spark.createDataFrame(data_2, schema=schema) df_filtered_mins_60 = df_filtered_mins_60.withColumn("start_timestamp", to_timestamp(df_filtered_mins_60["start_timestamp"], "yyyy-MM-dd'T'HH:mm:ss.SSS'+00:00'")) df_filtered_mins_60 = df_filtered_mins_60.withColumn("stop_timestamp", to_timestamp(df_filtered_mins_60["stop_timestamp"], "yyyy-MM-dd'T'HH:mm:ss.SSS'+00:00'")) display(df_filtered_mins_60)
解决方案
原代码的核心问题:
collect()将分布式数据拉取到Driver端本地处理,11万行数据会占用大量内存,且单线程遍历效率极低;- 每次循环调用
withColumn生成新DataFrame,导致数据血缘无限拉长,Spark优化器无法有效处理,计算量呈指数级增长。
正确的分布式处理代码如下:
from pyspark.sql import functions as F from pyspark.sql.window import Window # 为df_filtered_mins_60分配唯一自增counter df_filtered = df_filtered_mins_60.withColumn( "counter", F.row_number().over(Window.orderBy(F.monotonically_increasing_id())) ) # 关联两个DataFrame:按SKU匹配+时间窗口匹配 result_df = df_selected.join( df_filtered, (df_selected.CPSKU == df_filtered.CPSKU) & (df_selected.DATEUPDATED >= df_filtered.start_timestamp) & (df_selected.DATEUPDATED <= df_filtered.stop_timestamp), "left" # 保留df_selected中未匹配的行 ).select( df_selected["*"], F.coalesce(df_filtered.counter, F.lit(0)).alias("counter") # 未匹配行counter设为0,与原逻辑一致 ) display(result_df)
代码说明
- 分布式分配counter:用
row_number()结合monotonically_increasing_id()生成唯一自增ID,全程在分布式集群中处理,替代本地循环的i自增; - 批量关联匹配:通过
join一次性完成所有SKU和时间窗口的匹配,Spark会自动优化关联逻辑,利用集群资源高效处理千万级数据; - 兼容未匹配行:用
coalesce将未匹配行的counter设为0,与原代码逻辑保持一致。
该方案避免了本地循环和重复创建DataFrame,性能会有数量级提升,可在Databricks上高效运行。
内容的提问来源于stack exchange,提问作者irum zahra
相关产品推荐
相关产品推荐

