PySpark:如何更简洁地过滤结构体数组中的无效元素?
简化PySpark数组结构体元素过滤的实现方案
我需要处理一个PySpark DataFrame,其中price_history列是包含date、identifier、price字段的结构体数组。要求移除数组中满足以下任一条件的结构体元素:
date字段为nullprice字段为nullprice字段的值为0.00
数据定义与初始代码
from decimal import Decimal from datetime import date from pyspark.sql.functions import (array, when, array_except, col, lit, array_size, struct, max as pyspark_max) from pyspark.sql.types import (LongType, StringType, StructField, StructType, IntegerType, DateType, DecimalType, ArrayType) INPUT_SCHEMA = StructType([ StructField('ID', LongType(), True), StructField('brand', StringType(), True), StructField('price_history', ArrayType(StructType([ StructField('date', DateType(), True), StructField('identifier', StringType(), True), StructField('price', DecimalType(12, 2), True),]), True), True), ]) INPUT_DATA = [ [7572753287, 'brand 1', [[None, None, Decimal(2.19)], [date(2023, 2, 27), None, Decimal(1.79)]]], [7373874387383, 'brand 2', [[None, "N", Decimal(7.00)]]], [223278687678, 'brand 3', [[None, "NB", Decimal(0.63)]]], [2782872453, 'brand 2', [[None, "N", Decimal(2.19)]]], [86943438343, 'brand x', [[None, "NW", Decimal(1.49)]]], [87334273838, 'mybrand', [[date(2022, 3, 26), "NW", None], [date(2024, 1, 15), "E", Decimal(0.99)], [date(2024, 1, 15), "E", Decimal(0.00)]]], [2783783972, 'other brand', [[date(2024, 1, 15), "NW", None]]], ] EXPECTED_DATA = [ [7572753287, 'brand 1', [[date(2023, 2, 27), None, Decimal(1.79)]]], [7373874387383, 'brand 2', []], [223278687678, 'brand 3', []], [2782872453, 'brand 2', []], [86943438343, 'brand x', []], [87334273838, 'mybrand', [[date(2024, 1, 15), "E", Decimal(0.99)]]], [2783783972, 'other brand', []], ] expected_df = spark.createDataFrame(EXPECTED_DATA, INPUT_SCHEMA) input_df = spark.createDataFrame(INPUT_DATA, INPUT_SCHEMA) input_df.display()
原方案的问题
我之前实现的方案逻辑较为繁琐,需要先计算数组的最大长度,再循环处理每个元素,最后通过array_except移除构造的空结构体:
max_array_size = input_df.select(pyspark_max(array_size(col('price_history')))).collect()[0][0] empty_struct = struct(lit(None).cast(DateType()).alias('date'), lit(None).cast(StringType()).alias('identifier'), lit(None).cast(DecimalType(12, 2)).alias('price')) result = input_df.withColumn('price_history', array_except(array(*[when( (col('price_history').getItem(x).getField('date').isNull()) | (col('price_history').getItem(x).getField('price').isNull()) | (col('price_history').getItem(x).getField('price')==Decimal(0.00)), empty_struct).otherwise(col('price_history').getItem(x)) for x in range(max_array_size)]), array(empty_struct))) result.display()
更简洁的实现方式
在PySpark 3.1及以上版本中,可以直接使用filter函数对数组元素进行筛选,无需复杂的循环和空结构体构造:
from pyspark.sql.functions import filter, col # 筛选出date非空、price非空且price不等于0.00的结构体元素 filtered_df = input_df.withColumn( "price_history", filter( col("price_history"), lambda elem: elem["date"].isNotNull() & elem["price"].isNotNull() & (elem["price"] != Decimal(0.00)) ) ) filtered_df.display()
结果验证
使用以下代码验证结果是否符合预期(需PySpark 3.5及PyArrow支持):
from pyspark.testing import assertDataFrameEqual assertDataFrameEqual(expected_df, filtered_df)
这个方案逻辑直观,直接通过lambda表达式定义过滤规则,对数组中的每个结构体元素进行判断,保留符合要求的元素,代码量更少且可读性更高。
内容的提问来源于stack exchange,提问作者the_economist
相关产品推荐
相关产品推荐

