Spark中使用UDF依据数组列拆分字符串列的方案咨询
解决方案
最优方案:用Spark内置高阶函数实现(无需UDF,无循环,性能最优)
该方案全程在executor端分布式执行,不会拉取数据到driver,不需要编写for循环、列表推导,也无需调用collect()、自定义select逻辑,完全满足你的要求。
实现逻辑是通过aggregate高阶函数遍历长度数组,累积计算区间起始位置并截取对应字符串,直接得到拆分后的句子数组:
from pyspark.sql import functions as F result_df = df.withColumn( "split_sentences", F.aggregate( F.col("col2"), # 初始状态:起始位置0,空结果数组 F.struct(F.lit(0).alias("start_pos"), F.array().alias("sentences")), # 遍历每个长度值的处理逻辑 lambda acc, length: F.struct( (acc.start_pos + length).alias("start_pos"), # 注:Spark substring起始位置从1开始计数,因此需要+1 F.array_union( acc.sentences, F.array(F.substring(F.col("col1"), acc.start_pos + 1, length)) ).alias("sentences") ) ).sentences )
如果需要输出map结构(key为句子序号,value为句子内容),可以追加一步transform+map_from_entries处理:
result_df = result_df.withColumn( "sentence_map", F.map_from_entries( F.transform( F.col("split_sentences"), lambda sentence, idx: F.struct(F.concat(F.lit("句子_"), idx + 1), sentence) ) ) )
备选方案:自定义UDF实现
如果后续有更复杂的自定义拆分逻辑,可以用UDF实现,UDF的执行逻辑也全部在executor端,不会拉取数据到driver,可保证并行性:
from pyspark.sql import functions as F from pyspark.sql.types import ArrayType, StringType @F.udf(returnType=ArrayType(StringType())) def split_by_lengths(text, lengths): res = [] start = 0 for l in lengths: res.append(text[start:start+l]) start += l return res result_df = df.withColumn("split_sentences", split_by_lengths(F.col("col1"), F.col("col2")))
注:该方案UDF内部的循环仅在executor端的单条数据维度执行,不属于driver端全局循环,不会影响并行性。如果要求完全不使用任何for循环,优先选择第一种内置高阶函数方案。
内容的提问来源于stack exchange,提问作者dcrowley01
相关产品推荐
相关产品推荐

