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

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.10 15:28:14