PySpark DataFrame中移除长度小于阈值的数组末尾元素实现方法
PySpark 仅移除数组末尾长度不足阈值的元素实现
现有输入结构
给定包含id、text两个字段的PySpark DataFrame,构造代码与初始数据如下:
columns = ['id', 'text'] vals = [ (1, 'I am good or'), (2, 'You are okey'), (3, 'She is fine in') ] df = spark.createDataFrame(vals, columns)
初始数据预览:
+---+--------------+ | id| text| +---+--------------+ | 1| I am good or| | 2| You are okey| | 3|She is fine in| +---+--------------+
需求说明
将text字段按空格拆分为分词数组后,仅校验数组最后一个元素的长度:如果最后一个元素长度小于自定义阈值(示例阈值为3),则移除该末尾元素;其余位置的元素无论长度多少都保留,最终将处理后的分词数组拼接回文本字段。
原有错误写法的问题:使用filter遍历全数组元素,会删除所有长度小于阈值的元素,不符合仅处理末尾元素的要求,错误代码如下:
df.withColumn('split_text', F.split(F.col("text")," "))\ .withColumn("split_text", F.expr("filter(split_text, x -> not(length(x) < 3))"))
预期输出结果:
+---+------------+ | id| text| +---+------------+ | 1| I am good| | 2|You are okey| | 3| She is fine| +---+------------+
正确实现代码
核心逻辑是只对数组最后一位做长度判断,符合移除条件时用切片截掉最后一位,否则保留原数组,阈值可通过变量自定义:
import pyspark.sql.functions as F # 自定义末尾词最小长度阈值 MIN_LAST_WORD_LENGTH = 3 result = df.withColumn("split_text", F.split("text", " ")) \ .withColumn( "split_text", # 判断最后一个元素长度是否达标 F.when( F.length(F.element_at("split_text", -1)) < MIN_LAST_WORD_LENGTH, # 不达标则截取除最后一位外的所有元素 F.slice("split_text", 1, F.size("split_text") - 1) ).otherwise(F.col("split_text")) ) \ # 拼接回文本字段 .withColumn("text", F.concat_ws(" ", "split_text")) \ .drop("split_text")
关键逻辑说明
F.element_at(col, -1):直接取数组的最后一个元素,Spark 2.4及以上版本支持负数索引倒序取数的规则F.slice(col, start, length):从数组的start位置开始,截取length长度的元素,这里截取长度为数组总长度减1,刚好去掉最后一位不符合要求的元素- 全程没有遍历全数组做过滤,不会误删数组其他位置的短长度词汇
内容的提问来源于stack exchange,提问作者Rory
相关产品推荐
相关产品推荐

