如何使用Pyspark与shapely通过经纬度信息确定点位所属国家
点面匹配性能优化方案
现存代码核心性能瓶颈
- 重复WKT反序列化:现有代码每次调用UDF都会对所有国家的WKT字符串执行
wkt.loads,固定的地理数据重复计算产生了极大的冗余开销 - 暴力遍历匹配:每个点位都要遍历所有国家多边形执行
contains判断,时间复杂度为O(N)(N为国家数量),数据量上来后开销会线性增长 - 行式UDF的低效率:默认Pyspark行式UDF每次仅处理单条数据,序列化反序列化开销极高
具体优化实现
步骤1:提前预处理国家地理数据并广播
在Driver端提前完成WKT反序列化、空间索引构建,再将固定地理数据广播到所有Worker节点,避免重复计算:
import shapely.wkt as wkt from shapely.strtree import STRtree from shapely.geometry import Point from pyspark.sql import functions as f from pyspark.sql.functions import pandas_udf from pyspark.sql.types import StringType import pandas as pd # Driver端提前预处理国家地理数据 geometries = [] iso3_mapping = [] for geometry in geodata_countries_geo: poly = wkt.loads(geometry[0]) geometries.append(poly) iso3_mapping.append(geodata_countries_geo_dict[str(geometry[0])]) # 构建STR空间索引,快速过滤候选多边形 country_index = STRtree(geometries) # 广播预处理后的数据到所有Worker节点 broadcast_geo = spark.sparkContext.broadcast({ "index": country_index, "geoms": geometries, "iso3": iso3_mapping })
步骤2:使用Pandas UDF做批量匹配
替换行式UDF为批量处理的Pandas UDF,结合空间索引过滤候选多边形,大幅降低计算开销:
@pandas_udf(StringType()) def batch_match_country(lat_series: pd.Series, lon_series: pd.Series) -> pd.Series: geo_data = broadcast_geo.value index = geo_data["index"] geoms = geo_data["geoms"] iso3_list = geo_data["iso3"] result = [] for lat, lon in zip(lat_series, lon_series): point = Point(lon, lat) # 先通过空间索引拿到可能匹配的候选多边形,仅对候选做精确contains判断 candidate_idx = index.query(point) match_iso3 = None for idx in candidate_idx: if geoms[idx].contains(point): match_iso3 = iso3_list[idx] break result.append(match_iso3) return pd.Series(result) # 调用UDF完成匹配 dist_base = dist_base.withColumn('Country_Base_ISO3', batch_match_country(f.col("lat"),f.col("lon")))
额外优化建议
- 如果存在大量点位落在公海等无匹配区域,可以提前用所有国家的总边界框过滤无效点位,避免无意义的索引查询
- 若国家多边形精度过高,可以提前用
shapely.Object.simplify()做几何简化,降低contains判断的计算量,注意控制简化公差避免边界匹配出错 - 可以先对点位做经纬度分桶,再匹配对应区域的多边形,进一步降低单次匹配的候选数量
内容的提问来源于stack exchange,提问作者Alex Germain
相关产品推荐
相关产品推荐

