使用PandasUDFType.SCALAR处理Row数组报错:NumPy类型转换异常
解决PyArrow类型不匹配的问题
这个报错我之前踩过坑,本质是Spark和Pandas/PyArrow之间的整数类型映射不兼容导致的:Spark的IntegerType()对应32位整数,但Pandas在处理整数运算时默认会用64位整数(NumPy的int64),PyArrow没法自动把int64降级成int32,所以就抛出了这个ArrowInvalid错误。
咱们来修改你的代码,两种方案都可以选:
方案一:强制转换为int32类型(保留IntegerType返回)
如果必须要返回ArrayType(IntegerType()),可以在UDF里把处理后的元素强制转成32位整数:
import numpy as np from pyspark.sql import SparkSession from pyspark.sql.functions import pandas_udf from pyspark.sql.types import ArrayType, IntegerType spark = SparkSession.builder.getOrCreate() df = spark.createDataFrame([([1, 2, 3, 2],), ([4, 5, 5, 4],)], ['data']) @pandas_udf(ArrayType(IntegerType()), PandasUDFType.SCALAR) def s(x): # 这里注意:如果你的需求是**每个元素乘2**,要写成下面的方式 # 原来的xx*2是把数组重复一遍,不是元素乘2哦! z = x.apply(lambda xx: (np.array(xx) * 2).astype(np.int32).tolist()) # 如果你的需求是**重复数组两次**,就改成: # z = x.apply(lambda xx: (np.array(xx).repeat(2)).astype(np.int32).tolist()) return z df.select(s(df.data)).show()
方案二:改用LongType返回类型
如果业务允许用64位整数,直接把返回类型改成ArrayType(LongType()),这样就不用转换类型了:
from pyspark.sql.types import ArrayType, LongType @pandas_udf(ArrayType(LongType()), PandasUDFType.SCALAR) def s(x): # 不管是元素乘2还是数组重复,直接返回即可 z = x.apply(lambda xx: xx*2) return z
补充说明
你原来的lambda xx: xx*2是把列表重复两次(比如[1,2]变成[1,2,1,2]),如果你的真实需求是给每个元素乘2,一定要改成[i*2 for i in xx]或者用NumPy数组处理,不然逻辑会错哦!
内容的提问来源于stack exchange,提问作者littlely
相关产品推荐
相关产品推荐

