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

无需Spark UI,Databricks10.4下用PySpark检测Spark应用内存/磁盘溢出

在Databricks 10.4 PySpark中实现溢出(Spill)检测与自动化通知

核心思路

基于Spark内置的SparkListener扩展实现自定义溢出监听,通过onTaskEnd钩子捕获每个任务的溢出指标,设置阈值过滤"主要溢出",并触发自动化通知。


1. 自定义溢出监听类(SpillDetectionListener)

这个类会跟踪每个任务和Stage的溢出数据,根据预设阈值触发预警:

from pyspark.sparkcontext import SparkListener

class SpillDetectionListener(SparkListener):
    def __init__(self, mem_spill_threshold_mb=1024, disk_spill_threshold_mb=5120, spill_task_ratio_threshold=0.3):
        # 初始化阈值:单任务内存溢出1GB、磁盘溢出5GB;Stage溢出任务占比30%
        self.mem_threshold = mem_spill_threshold_mb * 1024 * 1024  # 转成字节单位
        self.disk_threshold = disk_spill_threshold_mb * 1024 * 1024
        self.spill_ratio_threshold = spill_task_ratio_threshold
        # 存储Stage级别的溢出统计数据
        self.stage_stats = {}

    def onTaskEnd(self, task_end_event):
        task_info = task_end_event.taskInfo
        task_metrics = task_end_event.taskMetrics
        stage_key = (task_info.stageId, task_info.stageAttemptId)

        # 提取当前任务的溢出数据
        mem_spill = task_metrics.memoryBytesSpilled
        disk_spill = task_metrics.diskBytesSpilled

        # 更新Stage统计数据
        if stage_key not in self.stage_stats:
            self.stage_stats[stage_key] = {
                "total_tasks": 0,
                "spilled_tasks": 0,
                "total_mem_spill": 0,
                "total_disk_spill": 0
            }
        stage_data = self.stage_stats[stage_key]
        stage_data["total_tasks"] += 1
        if mem_spill > 0 or disk_spill > 0:
            stage_data["spilled_tasks"] += 1
            stage_data["total_mem_spill"] += mem_spill
            stage_data["total_disk_spill"] += disk_spill

        # 检测单任务级主要溢出
        if mem_spill >= self.mem_threshold or disk_spill >= self.disk_threshold:
            self.send_alert(
                alert_type="任务级溢出",
                stage_id=task_info.stageId,
                attempt_id=task_info.stageAttemptId,
                task_id=task_info.taskId,
                mem_spill_mb=round(mem_spill / 1024 / 1024, 2),
                disk_spill_mb=round(disk_spill / 1024 / 1024, 2)
            )

        # 检测Stage级主要溢出(溢出占比或总溢出量超标)
        total_mem_mb = stage_data["total_mem_spill"] / 1024 / 1024
        total_disk_mb = stage_data["total_disk_spill"] / 1024 / 1024
        spill_ratio = stage_data["spilled_tasks"] / stage_data["total_tasks"] if stage_data["total_tasks"] > 0 else 0

        if (spill_ratio >= self.spill_ratio_threshold 
            or total_mem_mb >= self.mem_threshold / 1024 / 1024 * 2
            or total_disk_mb >= self.disk_threshold / 1024 / 1024 * 2):
            # 避免重复发送Stage预警,添加标记逻辑
            if "alert_sent" not in stage_data:
                self.send_alert(
                    alert_type="Stage级严重溢出",
                    stage_id=task_info.stageId,
                    attempt_id=task_info.stageAttemptId,
                    spill_ratio=round(spill_ratio * 100, 2),
                    total_mem_spill_mb=round(total_mem_mb, 2),
                    total_disk_spill_mb=round(total_disk_mb, 2),
                    total_tasks=stage_data["total_tasks"]
                )
                stage_data["alert_sent"] = True

    def send_alert(self, **kwargs):
        # 自定义通知逻辑,可替换为邮件、Teams、Slack等
        alert_msg = f"📢 {kwargs['alert_type']}预警\n"
        alert_msg += f"Stage ID: {kwargs['stage_id']}(Attempt {kwargs['attempt_id']})\n"
        if "task_id" in kwargs:
            alert_msg += f"任务ID: {kwargs['task_id']}\n"
            alert_msg += f"内存溢出: {kwargs['mem_spill_mb']} MB\n"
            alert_msg += f"磁盘溢出: {kwargs['disk_spill_mb']} MB\n"
        else:
            alert_msg += f"溢出任务占比: {kwargs['spill_ratio']}%\n"
            alert_msg += f"总内存溢出: {kwargs['total_mem_spill_mb']} MB\n"
            alert_msg += f"总磁盘溢出: {kwargs['total_disk_spill_mb']} MB\n"
            alert_msg += f"Stage总任务数: {kwargs['total_tasks']}\n"
        
        # 打印到控制台(Databricks集群日志可见)
        print(alert_msg)
        
        # 可选:Databricks邮件通知
        # import dbutils
        # dbutils.notifications.sendEmail(
        #     recipients=["your-team@example.com"],
        #     subject=f"Spark溢出预警 - {kwargs['alert_type']}",
        #     body=alert_msg
        # )

2. 注册监听类到SparkContext

在PySpark脚本开头注册Listener,确保当前会话生效:

from pyspark import SparkContext

# 获取当前SparkContext
sc = SparkContext.getOrCreate()

# 初始化监听类(可根据需求调整阈值)
spill_listener = SpillDetectionListener(
    mem_spill_threshold_mb=512,  # 单任务内存溢出512MB触发预警
    disk_spill_threshold_mb=2048,  # 单任务磁盘溢出2GB触发预警
    spill_task_ratio_threshold=0.2  # Stage中20%任务溢出触发预警
)

# 注册Listener
sc.addSparkListener(spill_listener)

3. 进阶优化

  • 持久化溢出数据:将stage_stats中的数据写入Delta表,用于后续性能分析:
    from pyspark.sql import Row
    from pyspark.sql import SparkSession
    
    spark = SparkSession.builder.getOrCreate()
    
    def save_spill_records():
        records = []
        for (stage_id, attempt_id), stats in spill_listener.stage_stats.items():
            records.append(Row(
                stage_id=stage_id,
                stage_attempt_id=attempt_id,
                total_tasks=stats["total_tasks"],
                spilled_tasks=stats["spilled_tasks"],
                total_mem_spill_mb=round(stats["total_mem_spill"]/1024/1024,2),
                total_disk_spill_mb=round(stats["total_disk_spill"]/1024/1024,2),
                spill_ratio=round(stats["spilled_tasks"]/stats["total_tasks"],4) if stats["total_tasks"]>0 else 0
            ))
        if records:
            spark.createDataFrame(records).write.mode("append").saveAsTable("monitoring.spill_analytics")
    
  • 集群级全局监听:将Listener注册逻辑放在集群初始化脚本(Init Script)中,实现所有会话自动监听。

注意事项

  • 阈值需根据任务规模和集群配置调整,避免误报或漏报。
  • Listener是会话级别的,重启SparkSession后需重新注册。
  • Databricks中使用dbutils.notifications需要确保集群有对应权限。

内容的提问来源于stack exchange,提问作者BI Dude

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.02 13:10:49