PySpark实现K-core图算法:迭代循环致Spark会话崩溃
PySpark实现K-core算法的心跳超时问题解决
背景与现有实现
我用PySpark实现K-core算法,初始数据准备、边统计关联以及迭代逻辑如下:
1. 初始数据准备(以k-core=10为例)
import networkx as nx G = nx.gnm_random_graph(100, 100) core = nx.k_core(G, 10) if len(core) > 0: df = pd.DataFrame(list(core.edges()), columns=['source', 'target']) graph = spark.createDataFrame(df)
2. 统计入边、出边并关联
outgoing_edge = graph.groupby('source').agg(F.collect_set('target').alias('count_outgoin')) ingoing_edge = graph.groupby('target').agg(F.collect_set('source').alias('count_ingoin')) ingoing_edge = ingoing_edge.withColumnRenamed("target", "source") get_all_count_edge = outgoing_edge.join(ingoing_edge, on="source", how="outer") \ .withColumn("count_outgoin", F.coalesce(F.col("count_outgoin"), F.array())) \ .withColumn("count_ingoin", F.coalesce(F.col("count_ingoin"), F.array())) \ .withColumn("in_plus_out", F.array_distinct(F.flatten(F.array(F.col('count_outgoin'), F.col('count_ingoin'))))) \ .withColumn('count_edge', F.size('in_plus_out')) \ .select('source', 'in_plus_out', 'count_edge') \ .cache() get_all_count_edge.count()
3. K-core迭代算法实现
def k_core_alg(get_all_count_edge): k = 2 dataframes = {} while True: print(f'=========== start k: {k} =========') get_all_count_edge_filter = get_all_count_edge.filter(F.col('count_edge') < k).select(F.col('source').alias('target')) print(f' see edge < k ') print(f'check count where edge < k: {get_all_count_edge_filter.count()}') graph = get_all_count_edge.filter(F.col('count_edge') >= k).withColumn('target',F.explode('in_plus_out')).select('source', 'target') print(f' see edge > k ') print(f'check count where edge > k (explode): {graph.count()}') graph = graph.join(F.broadcast(get_all_count_edge_filter), graph.target == get_all_count_edge_filter.target, 'left_anti') \ .cache() print(f' first join ') print(f'count first join: {graph.count()}') outgoing_edge = graph.groupby('source').agg(F.collect_set('target').alias('count_outgoin')) ingoing_edge = graph.groupby('target').agg(F.collect_set('source').alias('count_ingoin')) \ .withColumnRenamed("target", "source") get_all_count_edge = outgoing_edge.join(F.broadcast(ingoing_edge), on="source", how="outer") \ .withColumn("count_outgoin", F.coalesce(F.col("count_outgoin"), F.array())) \ .withColumn("count_ingoin", F.coalesce(F.col("count_ingoin"), F.array())) \ .withColumn("in_plus_out", F.array_distinct(F.flatten(F.array(F.col('count_outgoin'), F.col('count_ingoin'))))) \ .withColumn('count_edge', F.size('in_plus_out')) \ .select('source', 'in_plus_out', 'count_edge') print(f' second join ') print(f'second join: {get_all_count_edge.count()}') graph = get_all_count_edge.filter(F.col('count_edge') >= k).withColumn('target',F.explode('in_plus_out')).select('source', 'target') dataframes["prevgraph_" + str(k)] = graph count_check = graph.select('source').distinct().count() print(f' count distinct users: {count_check} ') if count_check <= 1: prev_k = k-1 k_core_graph_max = dataframes["prevgraph_" + str(prev_k)] break else: k += 1 k_end = prev_k print(f'k-core = {k_end}') return k_core_graph_max k_core_graph_max = k_core_alg(get_all_count_edge)
遇到的问题
随着k值增大,Spark处理时长逐渐增加,最终触发错误issue communicating with driver in heartbeater,导致Spark会话终止。要求完全基于PySpark实现,且保留打印语句用于定位崩溃节点。
优化方案
1. 主动释放缓存,避免内存泄漏
每次迭代生成的get_all_count_edge和graph会留下旧缓存,持续占用内存引发内存不足。在循环开始前添加缓存释放逻辑:
def k_core_alg(get_all_count_edge): k = 2 dataframes = {} prev_get_all_count_edge = None prev_graph = None while True: # 释放上一轮缓存 if prev_get_all_count_edge and prev_get_all_count_edge.is_cached: prev_get_all_count_edge.unpersist() if prev_graph and prev_graph.is_cached: prev_graph.unpersist() print(f'=========== start k: {k} =========') # ... 原有代码 ... # 更新缓存标记 prev_get_all_count_edge = get_all_count_edge prev_graph = graph
2. 减少不必要的count()调用
count()会触发全量计算,多次调用重复消耗资源:
- 可移除非关键的count打印,比如
check count where edge > k (explode); - 必须保留的count,先缓存DataFrame再调用,避免重复计算:
get_all_count_edge_filter = get_all_count_edge.filter(F.col('count_edge') < k).select(F.col('source').alias('target')).cache() print(f'check count where edge < k: {get_all_count_edge_filter.count()}')
3. 替换大数组操作,改用度数统计
原有逻辑用collect_set存储邻居节点,生成的大数组占用大量内存,改成直接统计节点邻居数:
def update_node_degree(graph): # 转成无向边 undirected_graph = graph.select('source', 'target').union(graph.select('target', 'source')) # 统计每个节点的唯一邻居数 node_degree = undirected_graph.groupBy('source').agg(F.countDistinct('target').alias('count_edge')) # 补充孤立节点 all_nodes = graph.select('source').union(graph.select('target')).distinct().alias('source') return all_nodes.join(node_degree, on='source', how='left') \ .fillna(0, subset=['count_edge'])
在迭代中用此函数替换原有的入边出边关联逻辑,大幅降低内存占用。
4. 调整Spark心跳与内存配置
计算耗时过长会触发心跳超时,修改配置延长超时时间并调整内存:
from pyspark.sql import SparkSession spark = SparkSession.builder \ .config("spark.network.timeout", "300s") # 全局网络超时 .config("spark.executor.heartbeatInterval", "60s") # executor心跳间隔 .config("spark.driver.memory", "16g") # 按集群情况调整 .config("spark.executor.memory", "16g") # 按集群情况调整 .getOrCreate()
5. 优化数据分区
避免单分区数据量过大,调整分区数提升并行度:
graph = graph.join(F.broadcast(get_all_count_edge_filter), graph.target == get_all_count_edge_filter.target, 'left_anti') \ .repartition(100) # 按集群节点数调整 .cache()
6. 动态判断广播使用
若get_all_count_edge_filter数据集较大,广播会占用过多executor内存,改用普通join:
filter_count = get_all_count_edge_filter.count() if filter_count < 10000: # 阈值按需调整 graph = graph.join(F.broadcast(get_all_count_edge_filter), graph.target == get_all_count_edge_filter.target, 'left_anti') else: graph = graph.join(get_all_count_edge_filter, graph.target == get_all_count_edge_filter.target, 'left_anti') graph = graph.cache()
内容的提问来源于stack exchange,提问作者Rory
相关产品推荐
相关产品推荐

