如何在PySpark GraphFrames中基于日期范围重叠规则定义边
问题描述
我正在使用PySpark的GraphFrames构建图,现有数据如下:
data = [ ("1990", "1995"), ("1980", "1996"), ("1993", "1994"), ("1990", "2002"), ("1996", "2002"), ("1999", "2008"), ("2003", "2014"), ]
节点由起始日期与结束日期拼接而成,示例节点如下:
nodes = [ "1990_1995", "1980_1996", "1993_1994", "1990_2002", "1996_2002", "1999_2008", "2003_2014", ]
边的定义规则为:若两个日期范围存在任何形式的重叠(包含、交叉等),则对应节点间创建一条边,期望边结果如下:
node_and_edges = [ ("1990_1995", ["1980_1996", "1993_1994", "1990_2002"]), ("1980_1996", ["1990_1995", "1993_1994", "1990_2002"]), ("1993_1994", ["1990_1995", "1980_1996", "1990_2002"]), ("1990_2002", ["1990_1995", "1980_1996", "1993_1994", "1996_2002", "1999_2008"]), ("1996_2002", ["1993_1994", "1990_2002", "1999_2008"]), ("1999_2008", ["1990_2002", "1996_2002", "2003_2014"]), ("2003_2014", ["1999_2008"]), ]
我已完成节点创建,现有代码如下:
from pyspark.sql import functions as F from pyspark.sql import SparkSession as ss spark = ss.builder.appName("test").getOrCreate() data = [ ("1990", "1995"), ("1980", "1996"), ("1993", "1994"), ("1990", "2002"), ("1996", "2002"), ("1999", "2008"), ("2003", "2014"), ] rdd = spark.sparkContext.parallelize(data) columns = ["start_date", "end_date"] df = rdd.toDF(columns) # Creating nodes df = df.withColumn("node", F.concat_ws("_", F.col("start_date"), F.col("end_date"))) # Creating edges # how? df.printSchema() df.show()
执行输出如下:
root |-- start_date: string (nullable = true) |-- end_date: string (nullable = true) |-- node: string (nullable = false) +----------+--------+---------+ |start_date|end_date| node| +----------+--------+---------+ | 1990| 1995|1990_1995| | 1980| 1996|1980_1996| | 1993| 1994|1993_1994| | 1990| 2002|1990_2002| | 1996| 2002|1996_2002| | 1999| 2008|1999_2008| | 2003| 2014|2003_2014| +----------+--------+---------+
请问如何在PySpark GraphFrames中实现上述边创建逻辑?
解决方案
1. 准备符合GraphFrames要求的节点DataFrame
GraphFrames要求节点必须包含id列作为唯一标识,直接转换现有节点列即可:
nodes_df = df.select(F.col("node").alias("id"))
2. 生成边DataFrame
通过自连接+重叠条件筛选实现边的创建:
# 将日期转为整数,避免字符串比较的逻辑错误 df_with_int_dates = df.withColumn("start_date_int", F.col("start_date").cast("int")) \ .withColumn("end_date_int", F.col("end_date").cast("int")) # 自连接匹配所有节点对,排除自身连接,筛选重叠的节点对 edges_df = df_with_int_dates.alias("a") \ .join(df_with_int_dates.alias("b"), on=F.col("a.node") != F.col("b.node")) \ # 核心重叠判断:a的起始早于b的结束,且a的结束晚于b的起始 .filter( (F.col("a.start_date_int") < F.col("b.end_date_int")) & (F.col("a.end_date_int") > F.col("b.start_date_int")) ) \ .select(F.col("a.node").alias("src"), F.col("b.node").alias("dst"))
3. 构建GraphFrame
导入GraphFrames并使用节点、边DataFrame创建图:
from graphframes import GraphFrame graph = GraphFrame(nodes_df, edges_df)
4. 验证结果(可选)
若要生成问题中期望的node_and_edges格式,可对边DataFrame进行聚合:
result_df = edges_df.groupBy("src") \ .agg(F.collect_list("dst").alias("edges")) \ .orderBy("src") result_df.show(truncate=False)
执行后输出与期望结果一致:
+---------+------------------------------------------------+ |src |edges | +---------+------------------------------------------------+ |1980_1996|[1990_1995, 1993_1994, 1990_2002] | |1990_1995|[1980_1996, 1993_1994, 1990_2002] | |1990_2002|[1990_1995, 1980_1996, 1993_1994, 1996_2002, 1999_2008]| |1993_1994|[1990_1995, 1980_1996, 1990_2002, 1996_2002] | |1996_2002|[1990_2002, 1993_1994, 1999_2008] | |1999_2008|[1990_2002, 1996_2002, 2003_2014] | |2003_2014|[1999_2008] | +---------+------------------------------------------------+
关键说明
- 日期转整数是为了避免字符串比较的逻辑错误(如
"200" > "1999"的错误判断) - 重叠条件
a.start < b.end AND a.end > b.start能覆盖所有类型的日期重叠场景 - 自连接时排除
a.node == b.node,避免节点自身创建无效边
内容的提问来源于stack exchange,提问作者shogitai
相关产品推荐
相关产品推荐

