PySpark中Haversine距离计算UDF执行Where子句时报错求助
解决PySpark中Haversine距离UDF过滤时的TypeError问题
嘿,我之前在把Python的Haversine函数转成PySpark UDF的时候也遇到过一模一样的问题!看起来计算结果正常,但一用WHERE过滤就报错,咱们来一步步搞定它:
先捋下你可能的代码场景
假设你的基础Python函数是这样的(和我当初写的差不多):
import math def haversine_distance(lat1, lon1, lat2, lon2): R = 6371 # 地球半径,单位公里 dlat = math.radians(lat2 - lat1) dlon = math.radians(lon2 - lon1) a = math.sin(dlat/2) * math.sin(dlat/2) + math.cos(math.radians(lat1)) * math.cos(math.radians(lat2)) * math.sin(dlon/2) * math.sin(dlon/2) c = 2 * math.atan2(math.sqrt(a), math.sqrt(1-a)) return R * c
然后你注册UDF的时候可能没指定返回类型,直接这么写了:
from pyspark.sql.functions import udf from pyspark.sql import SparkSession spark = SparkSession.builder.appName("HaversineTest").getOrCreate() # 这里没指定返回类型是坑的源头! haversine_udf = udf(haversine_distance) # 造点测试数据 data = [(40.7128, -74.0060, 34.0522, -118.2437), (51.5074, -0.1278, 48.8566, 2.3522)] df = spark.createDataFrame(data, ["lat1", "lon1", "lat2", "lon2"]) # 计算距离,看起来结果没问题 df = df.withColumn("dist_km", haversine_udf(df.lat1, df.lon1, df.lat2, df.lon2)) # 一过滤就炸了! df.where(df.dist_km > 1000).show()
为啥会报错?
这个TypeError主要有两个原因:
- UDF返回类型没明确指定:PySpark自动推断类型的时候经常会出错,比如把浮点型推断成其他类型,导致过滤时类型不匹配。
- 空值没处理:如果你的数据里有
null值,Python的math函数根本处理不了,直接就抛出异常了。
怎么解决?
方案1:给UDF明确指定返回类型
PySpark对UDF的类型要求很严格,必须明确告诉它返回的是什么类型,咱们改成这样:
from pyspark.sql.types import DoubleType # 加上DoubleType()指定返回类型,这一步至关重要! haversine_udf = udf(haversine_distance, DoubleType())
方案2:给函数加上空值判断
如果你的DataFrame里可能存在null值,一定要在函数里先做检查,不然一碰到空值就报错:
def haversine_distance(lat1, lon1, lat2, lon2): # 先检查有没有空值,有就返回None if None in (lat1, lon1, lat2, lon2): return None R = 6371 dlat = math.radians(lat2 - lat1) dlon = radians(lon2 - lon1) a = math.sin(dlat/2) **2 + math.cos(math.radians(lat1)) * math.cos(math.radians(lat2)) * math.sin(dlon/2)**2 c = 2 * math.atan2(math.sqrt(a), math.sqrt(1-a)) return R * c
方案3:直接用PySpark内置函数(推荐!)
其实完全没必要用Python UDF,PySpark有现成的数学函数,用这个方式性能更高,还不会有类型问题:
from pyspark.sql.functions import radians, sin, cos, sqrt, atan2 # 直接用PySpark的Column操作,不用写UDF def haversine_distance_spark(lat1, lon1, lat2, lon2): R = 6371 dlat = radians(lat2 - lat1) dlon = radians(lon2 - lon1) a = sin(dlat/2)**2 + cos(radians(lat1)) * cos(radians(lat2)) * sin(dlon/2)**2 c = 2 * atan2(sqrt(a), sqrt(1-a)) return R * c # 直接调用函数生成列 df = df.withColumn("dist_km", haversine_distance_spark(df.lat1, df.lon1, df.lat2, df.lon2)) df.where(df.dist_km > 1000).show()
验证一下
修改完之后再运行过滤,就正常了!比如上面的测试数据,纽约到洛杉矶的距离大概3940km,会被筛选出来,伦敦到巴黎的344km会被过滤掉。
内容的提问来源于stack exchange,提问作者E B
相关产品推荐
相关产品推荐

