如何在PySpark DataFrame中从含列表的现有列生成新列?
PySpark 从嵌套数组列生成新列的解决方案
先创建测试DataFrame
首先还原原始DataFrame:
from pyspark.sql import SparkSession from pyspark.sql.types import ArrayType, StructType, StructField, StringType spark = SparkSession.builder.appName("trans_parse").getOrCreate() # 定义数据结构 schema = StructType([ StructField("id", StringType(), nullable=False), StructField("transaction", ArrayType(ArrayType(StringType())), nullable=False) ]) # 测试数据 data = [ ("06t84g", [["T_BILL", "0.99"], ["Z_BILL", "0.33"], ["A_BILL", "0.77"]]), ("098t1g", [["T_BILL", "0.419"], ["Z_BILL", "0.19"], ["A_BILL", "0.137"]]), ("03z94f", [["T_BILL", "0.79"], ["Z_BILL", "0.49"], ["A_BILL", "0.317"]]), ("10yw22", [["T_BILL", "0.91"], ["Z_BILL", "0.818"], ["A_BILL", "0.457"]]), ("30r990", [["T_BILL", "0.193"], ["Z_BILL", "0.69"], ["A_BILL", "0.947"]]) ] df = spark.createDataFrame(data, schema=schema)
方法一:转换为Map类型后提取字段(已知所有列名)
如果明确知道要提取的账单类型,可将嵌套数组转为Map后直接提取对应值:
from pyspark.sql.functions import create_map, expr, col # 将transaction数组转为Map,键为账单类型,值转为float类型 df_with_map = df.withColumn( "trans_map", create_map( expr("transaction[0][0]"), expr("cast(transaction[0][1] as float)"), expr("transaction[1][0]"), expr("cast(transaction[1][1] as float)"), expr("transaction[2][0]"), expr("cast(transaction[2][1] as float)") ) ) # 提取目标列并移除临时map列 result_df = df_with_map.select( "id", "transaction", col("trans_map.T_BILL").alias("T_BILL"), col("trans_map.Z_BILL").alias("Z_BILL"), col("trans_map.A_BILL").alias("A_BILL") ).drop("trans_map") result_df.show()
优点:性能高效,使用PySpark内置函数,无数据shuffle。
方法二:Explode + Pivot(通用方案)
如果不确定所有账单类型或需扩展性,用explode拆分数组后pivot转列更通用:
from pyspark.sql.functions import explode, col # 拆分嵌套数组为单独行,提取账单类型和金额 exploded_df = df.select( "id", "transaction", explode("transaction").alias("trans_item") ).select( "id", "transaction", col("trans_item")[0].alias("bill_type"), col("trans_item")[1].cast("float").alias("amount") ) # 按id和transaction分组,将bill_type转为列 pivoted_df = exploded_df.groupBy("id", "transaction").pivot("bill_type").sum("amount") pivoted_df.show()
优点:自动识别所有账单类型,无需硬编码列名,扩展性强。
方法三:自定义UDF(不推荐)
若逻辑复杂内置函数无法满足,可使用UDF,但性能远低于内置函数,大数据量场景不建议:
from pyspark.sql.functions import udf from pyspark.sql.types import FloatType def get_bill_amount(trans_list, bill_name): for item in trans_list: if item[0] == bill_name: return float(item[1]) return None # 注册UDF extract_bill_udf = udf(get_bill_amount, FloatType()) # 提取目标列 result_df = df.select( "id", "transaction", extract_bill_udf(col("transaction"), "'T_BILL'").alias("T_BILL"), extract_bill_udf(col("transaction"), "'Z_BILL'").alias("Z_BILL"), extract_bill_udf(col("transaction"), "'A_BILL'").alias("A_BILL") ) result_df.show()
内容的提问来源于stack exchange,提问作者ASH
相关产品推荐
相关产品推荐

