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

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.14 13:17:02