You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.09.26 18:45:02