Databricks中Python多进程代码执行异常问题排查
问题分析与解决方案
核心问题
你在Databricks中用multiprocessing实现并行的思路存在几个致命问题,导致多进程方案失效:
- 函数嵌套错误:
consolidate_base函数被定义在update_campaing_headcounter_sum内部,外部调用时会找不到该函数,属于语法层面的错误。 - SparkSession进程安全问题:子进程无法继承Driver的SparkSession,
compute_new_base_values、get_cassandra_sum中调用的Cassandra读写操作依赖SparkSession,在子进程中会直接报错。 - 违背Spark分布式理念:把整个DataFrame用
collect()拉到Driver端再拆分,不仅容易触发内存溢出,还完全浪费了Spark集群的分布式计算能力——multiprocessing只能在Driver节点本地并行,无法利用Worker节点资源。
正确的并行实现方案
在Databricks/Spark环境中,应该利用Spark的分布式计算能力,而非本地多进程。以下是重构后的代码思路:
步骤1:修复函数结构
将嵌套的consolidate_base移到顶层,避免作用域错误。
步骤2:用Spark分布式操作替代本地多进程
使用foreachPartition让每个分区在Worker节点上并行处理,每个分区内批量处理数据,同时确保分区内的Spark操作能正确获取上下文。
重构后的代码示例
from pyspark.sql import DataFrame from pyspark.sql import functions as F from datetime import datetime, timedelta import time def get_hc_info_from_active_campaign() -> DataFrame: query = '''( select hc.id HeadcounterId ,hc.Name as HeadcounterName ,hc.[DeviceId] ,c.[id] as CampaignId ,c.[startDate] ,c.[stopDate] FROM [dbo].[MD_Headcounters] hc join MD_Campaigns as C on hc.CampaignId=C.Id where c.StopDate >CONVERT(date,GETDATE()-1)) as hc_info_from_active_campaign ''' hc_info_from_active_campaign = (spark.read .format("jdbc") .option("url", "jdbc:sqlserver://{0}:{1};database={2}".format(jdbcHostname_prod, jdbcPort, jdbcDatabase_prod)) .option("dbtable", query) .option("user", sql_user_prod) .option("password", sql_password_prod) .load() ) return hc_info_from_active_campaign def compute_new_base_values(timestamp_inf: str, timestamo_sup: str, device_id: str, category: str): cassandra_h_aggr_df = read_from_cassandra_prod('headcounter_category_h_aggr').where((F.col('device_id')== device_id) & (F.col('timestamp') <= timestamo_sup) & (F.col('timestamp') >= timestamp_inf) & (F.col('category')==category)).fillna(0) cassandra_h_aggr_consolidated_hystory_df = (cassandra_h_aggr_df .withColumn('delta_ots', F.greatest(F.col('estimated_ots'), F.col('ots')) - F.least(F.col('estimated_ots'), F.col('ots'))) .withColumn('delta_views', F.greatest(F.col('estimated_views'), F.col('views')) - F.least(F.col('estimated_views'), F.col('views'))) .withColumn('delta_attentiont_milliseconds', F.greatest(F.col('estimated_attentiont_milliseconds'), F.col('attentiont_milliseconds')) - F.least(F.col('estimated_attentiont_milliseconds'), F.col('attentiont_milliseconds'))) .withColumn('delta_dwell_milliseconds', F.greatest(F.col('estimated_dwell_milliseconds'), F.col('dwell_milliseconds')) - F.least(F.col('estimated_dwell_milliseconds'), F.col('dwell_milliseconds'))) .withColumn('min_ots', F.least(F.col('estimated_ots'), F.col('ots'))) .withColumn('min_views', F.least(F.col('estimated_views'), F.col('views'))) .withColumn('min_attentiont_milliseconds', F.least(F.col('estimated_attentiont_milliseconds'), F.col('attentiont_milliseconds'))) .withColumn('min_dwell_milliseconds', F.least(F.col('estimated_dwell_milliseconds'), F.col('dwell_milliseconds'))) ) consolidated_hystory_df = (cassandra_h_aggr_consolidated_hystory_df .withColumn('base_ots', F.col('delta_ots')+F.col('min_ots')) .withColumn('base_views', F.col('delta_views')+F.col('min_views')) .withColumn('base_attentiont_milliseconds', F.col('delta_attentiont_milliseconds')+F.col('min_attentiont_milliseconds')) .withColumn('base_dwell_milliseconds', F.col('delta_dwell_milliseconds')+F.col('min_dwell_milliseconds')) ) consolidated_hystory_df = (consolidated_hystory_df .groupBy('device_id', 'category') .agg( F.sum('estimated_ots').alias('sum_estimated_ots'), F.sum('base_ots').alias('sum_base_ots'), F.sum('ots').alias('sum_ots'), F.sum('estimated_views').alias('sum_estimated_views'), F.sum('base_views').alias('sum_base_views'), F.sum('views').alias('sum_views'), F.sum('estimated_attentiont_milliseconds').alias('sum_estimated_attentiont_milliseconds'), F.sum('base_attentiont_milliseconds').alias('sum_base_attentiont_milliseconds'), F.sum('attentiont_milliseconds').alias('sum_attentiont_milliseconds'), F.sum('estimated_dwell_milliseconds').alias('sum_estimated_dwell_milliseconds'), F.sum('base_dwell_milliseconds').alias('sum_base_dwell_milliseconds'), F.sum('dwell_milliseconds').alias('sum_dwell_milliseconds') ) ) updated_base_values_df =(consolidated_hystory_df .withColumn('BaseOts', F.col('sum_base_ots')-F.col('sum_ots')) .withColumn('BaseViews', F.col('sum_base_views')-F.col('sum_views')) .withColumn('BaseDwellTimeMilliseconds', F.col('sum_base_dwell_milliseconds')-F.col('sum_dwell_milliseconds')) .withColumn('BaseAttentionTimeMilliseconds', F.col('sum_base_attentiont_milliseconds')-F.col('sum_attentiont_milliseconds')) ) # 增加空数据判断,避免索引越界报错 if updated_base_values_df.count() > 0: return updated_base_values_df.select('BaseOts','BaseViews','BaseDwellTimeMilliseconds','BaseAttentionTimeMilliseconds').rdd.map(lambda row: row.asDict()).collect()[0] else: return {"BaseOts":0, "BaseViews":0, "BaseDwellTimeMilliseconds":0, "BaseAttentionTimeMilliseconds":0} def get_cassandra_sum(campaign_id: str, headcounter_id: str, category: str): df = read_from_cassandra_dev('temp_sum_headcounter_category').where((F.col('campaign_id')==campaign_id) & (F.col('headcounter_id')==headcounter_id) & (F.col('category')==category)) return df def process_partition(partition): # 每个分区内初始化时间参数 yesterday = datetime.now() - timedelta(days = 1) yesterday = yesterday.date() for elem in partition: timestamp_inf = elem['startDate'].strftime("%Y-%m-%d %H:00:00") timestamo_sup = yesterday.strftime("%Y-%m-%d 23:00:00") device_id = elem['DeviceId'] headcounter_id = elem['HeadcounterId'] campaign_id = elem['CampaignId'] for category in ['PERSON','VEHICLE']: new_base_values = compute_new_base_values(timestamp_inf, timestamo_sup, device_id, category) temp_sum_headcounter_category_df = get_cassandra_sum(campaign_id, headcounter_id, category) if temp_sum_headcounter_category_df.count() > 0: new_temp_headcounter_category_df = (temp_sum_headcounter_category_df .withColumn('base_ots',new_base_values['BaseOts']- F.col('base_ots')) .withColumn('base_views',new_base_values['BaseViews']- F.col('base_views')) .withColumn('base_dwell_milliseconds',new_base_values['BaseDwellTimeMilliseconds']- F.col('base_dwell_milliseconds')) .withColumn('base_attention_milliseconds',new_base_values['BaseAttentionTimeMilliseconds'] - F.col('base_attention_milliseconds')) ).select('campaign_id','headcounter_id','category','base_ots', 'base_views','base_dwell_milliseconds','base_attention_milliseconds') write_on_cassandra_dev(new_temp_headcounter_category_df, 'temp_sum_headcounter_category') def consolidate_base(hc_info_df): # 用Spark的foreachPartition实现分布式并行 hc_info_df.rdd.foreachPartition(process_partition) # 主逻辑 days_intervall = 1 start_script = time.time() hc_info_from_active_campaign_df = get_hc_info_from_active_campaign() print('--------') consolidate_base(hc_info_from_active_campaign_df) print(f"执行完成,耗时:{time.time() - start_script}秒")
关键优化点
- 移除本地多进程:用
foreachPartition替代multiprocessing,让数据在Worker节点分布式并行处理,充分利用集群资源。 - 避免数据拉到Driver:直接对DataFrame进行操作,不再用
collect()把数据拉到Driver端,避免内存溢出和性能损耗。 - 修复函数作用域:把嵌套函数移到顶层,解决调用错误。
- 增加空数据判断:在
compute_new_base_values中增加空DataFrame判断,避免索引越界报错。
额外建议
- 如果Cassandra的读写操作是性能瓶颈,可以考虑用Spark的Cassandra Connector批量读写,减少单条操作的开销。
- 尽量用Spark的DataFrame API替代RDD操作,提升优化效率。
内容的提问来源于stack exchange,提问作者Gabriele Sciurti
相关产品推荐
相关产品推荐

