PySpark自定义UDF处理分词如何返回ArrayType数组而非字符串
问题根因
你自定义的UDF没有显式声明返回值类型,PySpark默认会将Python函数的返回值序列化为StringType,因此得到的removed字段是字符串而非预期的数组类型。
方案1:修改UDF声明指定返回类型
仅需在注册UDF时明确指定返回类型为ArrayType(StringType())即可,修改后的代码如下:
from pyspark.sql.functions import udf, col from pyspark.sql.types import * def remove_stop_words(list_of_tokens, list_of_stopwords): ''' 接收分词列表和停用词列表,返回过滤停用词后的分词列表 ''' return [token for token in list_of_tokens if token not in list_of_stopwords] def udf_remove_stop_words(list_of_stopwords): ''' 生成绑定了指定停用词列表的UDF,显式声明返回类型为字符串数组 ''' # 新增第二个参数指定返回类型 return udf(lambda x: remove_stop_words(x, list_of_stopwords), ArrayType(StringType())) wordsNoStopDF = splitworddf.withColumn('removed', udf_remove_stop_words(list_of_words_to_get_rid)(col('split')))
修改完成后可执行wordsNoStopDF.printSchema()验证,removed字段类型会显示为array<string>,和split字段完全一致。
方案2:使用内置高阶函数(更优)
PySpark 2.4及以上版本支持内置高阶函数,性能远高于Python UDF,适合处理大数据量语料,且无需额外处理返回类型,实现逻辑如下:
from pyspark.sql.functions import expr, col stopwords = list_of_words_to_get_rid # 拼接停用词过滤的SQL表达式 filter_expr = f"filter(split, x -> x not in ({','.join([f'\'{w.replace(\"'\", \"''\")}\'' for w in stopwords])}))" wordsNoStopDF = splitworddf.withColumn('removed', expr(filter_expr))
注:上述代码中对停用词的单引号做了转义处理,避免停用词包含单引号时出现SQL语法错误;该方案完全保留原分词数组的重复词逻辑,和你自定义UDF的过滤效果一致。
修改完成后你可以直接对removed字段调用explode函数做后续的词频统计。
内容的提问来源于stack exchange,提问作者Michael
相关产品推荐
相关产品推荐

