如何用PySpark GraphFrame与NetworkX构建大洲-国家-城市层级图?
实现「大洲-国家-城市」层级图的正确方法
问题分析
你的代码存在几个核心问题:
- 顶点不完整:仅将国家作为顶点,遗漏了大洲和城市节点,层级关系的基础结构缺失
- 边构造错误:错误地将「大洲-国家」「国家-城市」两种层级关系合并为单条边的多列结构,不符合图结构中「源-目标」的二元边定义
- NetworkX转换参数错误:
nx.from_pandas_edgelist不支持同时传入三个节点列,无法正确识别层级关联逻辑
正确实现步骤
1. 读取并预处理数据(基于示例数据源结构)
先确保所有基础表和关联表正确读取:
from pyspark.sql import SparkSession from pyspark.sql.functions import col import networkx as nx import matplotlib.pyplot as plt # 初始化SparkSession(未初始化时执行) spark = SparkSession.builder.appName("ContinentCountryCityGraph").getOrCreate() # 读取基础表 city = spark.read.csv("city.csv", header='true').withColumnRenamed("name", "city_name") country = spark.read.csv("country.csv", header='true').withColumnRenamed("name", "country_name") continent = spark.read.csv("continent.csv", header='true').withColumnRenamed("name", "continent_name") # 读取关联表(大洲-国家、国家-城市的映射关系) country_continent = spark.read.csv("country_continent.csv", header='true') city_country = spark.read.csv("city_country.csv", header='true') # 对应原代码中的metro_country
2. 构造完整的顶点集合
层级图需要包含大洲、国家、城市三类节点,统一用id作为标识列:
# 提取各类节点作为顶点 continent_vertices = continent.select(col("continent_name").alias("id")) country_vertices = country.select(col("country_name").alias("id")) city_vertices = city.select(col("city_name").alias("id")) # 合并所有顶点并去重 all_vertices = continent_vertices.union(country_vertices).union(city_vertices).distinct()
3. 构造两类层级边
分别创建「大洲→国家」和「国家→城市」的二元边,统一用src(源节点)和dst(目标节点)命名:
# 构造大洲-国家的边:大洲为源,国家为目标 continent_country_edges = country.join(country_continent, country.country_id == country_continent.country_id)\ .join(continent, country_continent.continent_id == continent.continent_id)\ .select(col("continent_name").alias("src"), col("country_name").alias("dst")) # 构造国家-城市的边:国家为源,城市为目标 country_city_edges = country.join(city_country, country.country_id == city_country.country_id)\ .join(city, city_country.city_id == city.city_id)\ .select(col("country_name").alias("src"), col("city_name").alias("dst")) # 合并两类边 all_edges = continent_country_edges.union(country_city_edges)
4. 创建GraphFrame并转换为NetworkX图
# 创建完整层级图的GraphFrame full_graph = GraphFrame(all_vertices, all_edges) # 转换为NetworkX图 nx_full_graph = nx.from_pandas_edgelist(full_graph.edges.toPandas(), 'src', 'dst') # 绘制完整层级图(调整参数优化显示) plt.figure(figsize=(12, 8)) nx.draw(nx_full_graph, with_labels=True, node_size=300, font_size=10, edge_color="#2ecc71", node_color="#3498db") plt.title("Continent -> Country -> City Hierarchy Graph") plt.show()
5. 过滤特定大洲的子图(以北美洲为例)
如果仅需展示单个大洲的层级结构,直接过滤关联边即可:
# 获取北美洲的所有国家列表 na_countries = country.join(country_continent, country.country_id == country_continent.country_id)\ .join(continent, country_continent.continent_id == continent.continent_id)\ .filter(col("continent_name") == "North America")\ .select("country_name").rdd.flatMap(lambda x: x).collect() # 过滤北美洲相关的所有边(包含大洲→国家、国家→城市) north_america_edges = all_edges.filter( (col("src") == "North America") | (col("src").isin(na_countries)) ) # 提取子图的顶点(从过滤后的边中获取所有节点) north_america_vertices = north_america_edges.select("src").union(north_america_edges.select("dst")).distinct().withColumnRenamed("src", "id") # 创建北美洲子图 north_america_graph = GraphFrame(north_america_vertices, north_america_edges) # 转换为NetworkX并绘制 nx_na_graph = nx.from_pandas_edgelist(north_america_graph.edges.toPandas(), 'src', 'dst') plt.figure(figsize=(10, 6)) nx.draw(nx_na_graph, with_labels=True, node_size=400, font_size=10, edge_color="#e74c3c", node_color="#f39c12") plt.title("North America -> Country -> City Hierarchy") plt.show()
关键说明
- 所有层级节点必须纳入顶点集合,否则绘制时会缺失顶端(大洲)或底端(城市)节点
- 边必须是二元结构(源→目标),两类层级关系需分开构造后合并,不能混在同一行中
- 过滤子图时,需确保包含目标大洲的所有关联节点,避免层级断裂
内容的提问来源于stack exchange,提问作者Kaykay38
相关产品推荐
相关产品推荐

