如何在PySpark中提取数组列中指定长度的字符串元素?
在PySpark DataFrame中提取数组列里长度为2的字符串元素
问题描述
我有一个包含数组列的PySpark DataFrame,该数组列的每个元素均为字符串,需要提取其中长度为2的元素。
输入数据示例:
| id | array_with_strings |
|---|---|
| 00001 | [N, NS, NSY, NSB] |
| 00002 | [B, BS, BSNY, BSD] |
期望输出:
| id | twoDigits |
|---|---|
| 00001 | NS |
| 00002 | BS |
我尝试了以下代码,但不确定是否正确:
def getTwoDigits(arr): for x in arr: if F.length(x) == 2: return x else: return None extractTwoDigits_udf = F.udf(lambda z: getTwoDigits(z)) df = df.withColumn( "twoDigits", extractTwoDigits_udf(F.col("array_with_strings")) )
问题分析与解决方法
1. 现有UDF的问题
你写的UDF存在逻辑错误:循环中第一个元素不满足长度为2的条件时,会直接返回None,不会继续遍历数组后续元素。比如第一行数组的第一个元素是"N"(长度1),函数直接返回None,根本不会检查到后面的"NS"。
2. 修复后的UDF方法
如果坚持用UDF,需要调整逻辑,遍历完整个数组后再返回默认值,同时注意UDF内部使用Python原生的len()而非PySpark的F.length():
from pyspark.sql import functions as F def getTwoDigits(arr): for x in arr: if len(x) == 2: return x return None # 遍历完无匹配元素时返回None extractTwoDigits_udf = F.udf(getTwoDigits) df = df.withColumn("twoDigits", extractTwoDigits_udf(F.col("array_with_strings")))
3. 推荐:使用PySpark内置函数(无需UDF)
PySpark内置函数的性能远优于自定义UDF,推荐用filter+element_at组合实现需求:
from pyspark.sql import functions as F df = df.withColumn( "twoDigits", F.element_at( # 过滤数组中长度为2的字符串 F.filter(F.col("array_with_strings"), lambda x: F.length(x) == 2), 1 # 取过滤后数组的第一个元素,无匹配则返回null ) )
测试验证
用示例数据测试上述内置函数方法:
# 构造示例DataFrame data = [("00001", ["N", "NS", "NSY", "NSB"]), ("00002", ["B", "BS", "BSNY", "BSD"])] df = spark.createDataFrame(data, ["id", "array_with_strings"]) # 应用处理逻辑 df = df.withColumn( "twoDigits", F.element_at(F.filter(F.col("array_with_strings"), lambda x: F.length(x) == 2), 1) ) # 查看结果 df.show()
输出结果:
+-----+-----------------+---------+ | id|array_with_strings|twoDigits| +-----+-----------------+---------+ |00001| [N, NS, NSY, NSB]| NS| |00002| [B, BS, BSNY, BSD]| BS| +-----+-----------------+---------+
完全符合期望输出。
内容的提问来源于stack exchange,提问作者Cam
相关产品推荐
相关产品推荐

