PySpark如何获取array列所有True元素索引并提取对应源数组值
问题核心原因
- 原UDF使用
index("TRUE")仅能获取第一个匹配项的索引,无法遍历所有位置收集符合条件的元素 - 存在变量名错误:未定义
country_value变量,直接使用输入参数country即可 - 存在类型不匹配隐患:UDF声明返回
ArrayType,但country非空时返回字符串,会触发Spark列类型不一致报错
方案1:修改自定义UDF实现
Spark要求同一列数据类型必须一致,因此优先统一列类型为ArrayType(StringType),非空country也包装为单元素数组,若需要字符串格式可后续单独转换:
from pyspark.sql.types import StringType, ArrayType import pyspark.sql.functions as F @udf(returnType=ArrayType(StringType())) def determine_entity_country(country: str, sources: list, infer_from_source: list) -> list: if country is not None: return [country] res = [] # 遍历所有索引收集符合条件的source元素 for idx, is_true in enumerate(infer_from_source): if is_true == "TRUE": res.append(sources[idx]) return res if res else None # 调用UDF生成新列 df = df.withColumn("inferred_country", determine_entity_country(F.col("country"), F.col("sources"), F.col("infer_from_source")))
如果仅用于结果输出、需要和你给出的预期格式完全一致,可将UDF返回类型改为StringType(),调整返回逻辑即可:
@udf(returnType=StringType()) def determine_entity_country(country: str, sources: list, infer_from_source: list) -> str: if country is not None: return country res = [] for idx, is_true in enumerate(infer_from_source): if is_true == "TRUE": res.append(sources[idx]) return str(res) if res else None
方案2:使用内置函数实现(性能更优,推荐大数据量场景使用)
不需要自定义UDF,直接用PySpark原生数组函数处理,避免UDF的序列化开销:
df = df.withColumn( "inferred_country", F.when( F.col("country").isNotNull(), F.array(F.col("country")) ).otherwise( # 按位置配对两个数组→过滤TRUE项→提取对应source元素 F.transform( F.filter( F.arrays_zip(F.col("sources"), F.col("infer_from_source")), lambda x: x.infer_from_source == "TRUE" ), lambda x: x.sources ) ) )
内容的提问来源于stack exchange,提问作者Tytire Recubans
相关产品推荐
相关产品推荐

