无需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
相关产品推荐
相关产品推荐

