PySpark扁平化DataFrame时追加父列名称的实现方法
在PySpark中扁平化DataFrame并将父列名追加到子列名
嘿,这个需求我之前也遇到过,其实通过递归处理DataFrame的嵌套Schema就能轻松实现!我给你一步步拆解,先从示例场景开始,再给出通用的解决方案。
1. 先定义一个示例嵌套DataFrame
首先我们创建一个带有嵌套Struct和Array的DataFrame,方便演示:
from pyspark.sql import SparkSession from pyspark.sql.types import StructType, StructField, StringType, IntegerType, ArrayType spark = SparkSession.builder.appName("FlattenWithParentPrefix").getOrCreate() # 定义嵌套Schema sample_schema = StructType([ StructField("id", IntegerType(), nullable=True), StructField("user_profile", StructType([ StructField("full_name", StringType(), nullable=True), StructField("age", IntegerType(), nullable=True), StructField("contact", StructType([ StructField("email", StringType(), nullable=True), StructField("phone", StringType(), nullable=True) ])) ])), StructField("recent_orders", ArrayType(StructType([ StructField("order_id", StringType(), nullable=True), StructField("total", IntegerType(), nullable=True) ]))) ]) # 测试数据 sample_data = [ (1, ("Alice Smith", 32, ("alice@example.com", "123-456-7890")), [("ORD001", 99), ("ORD002", 199)]), (2, ("Bob Johnson", 28, ("bob@example.com", "987-654-3210")), [("ORD003", 49)]) ] df = spark.createDataFrame(sample_data, schema=sample_schema) df.printSchema()
运行后你会看到嵌套的Schema结构,接下来我们要把它扁平化,同时把父列名拼到子列名前面(比如user_profile.full_name变成user_profile_full_name)。
2. 实现通用的扁平化函数
核心思路是递归遍历Schema的每个字段:
- 如果是
StructType,就递归处理它的子字段,拼接父列名和子列名 - 如果是
ArrayType,可以选择explode展开数组(或者保留数组结构,看需求),再处理里面的嵌套字段 - 简单类型直接重命名(如果有父列前缀的话)
方案一:展开数组并完全扁平化
如果你需要把数组展开成多行,同时扁平化所有嵌套字段,可以用这个函数:
from pyspark.sql.functions import col, explode from pyspark.sql.types import StructType, ArrayType def flatten_with_parent_prefix(df, parent_prefix=""): selected_fields = [] for field in df.schema.fields: # 拼接父列前缀和当前列名,用下划线分隔(可自定义分隔符) new_col_name = f"{parent_prefix}_{field.name}" if parent_prefix else field.name if isinstance(field.dataType, StructType): # 递归处理嵌套Struct:先把当前Struct列重命名,再展开它的子字段 nested_df = df.select(col(field.name).alias(new_col_name)).select(f"{new_col_name}.*") # 递归调用后得到扁平化的子字段,添加到列表中 selected_fields.extend(flatten_with_parent_prefix(nested_df, new_col_name).columns) elif isinstance(field.dataType, ArrayType): # 先explode数组,处理里面的元素 exploded_df = df.withColumn(f"exploded_{field.name}", explode(col(field.name))) if isinstance(field.dataType.elementType, StructType): # 数组元素是Struct,继续递归扁平化 nested_df = exploded_df.select(col(f"exploded_{field.name}").alias(new_col_name)).select(f"{new_col_name}.*") flattened_nested = flatten_with_parent_prefix(nested_df, new_col_name) # 保留原DataFrame中非数组的列,加上扁平化后的数组子列 original_non_array_cols = [col(c) for c in df.columns if c != field.name] selected_fields = original_non_array_cols + flattened_nested.columns else: # 数组元素是简单类型,直接重命名 selected_fields.append(col(f"exploded_{field.name}").alias(new_col_name)) else: # 简单数据类型,直接添加(如果有前缀就重命名) selected_fields.append(col(field.name).alias(new_col_name)) return df.select(selected_fields)
测试这个函数
flattened_df = flatten_with_parent_prefix(df) flattened_df.printSchema() flattened_df.show(truncate=False)
你会看到输出的列名都是父列名_子列名的格式,比如user_profile_full_name、user_profile_contact_email,数组也被展开成了多行。
方案二:保留数组结构,仅扁平化数组内的Struct
如果你不想展开数组,而是保留数组结构但扁平化里面的Struct元素,可以用这个版本:
from pyspark.sql.functions import col, struct, transform from pyspark.sql.types import StructType, ArrayType def flatten_array_structs_with_prefix(df, parent_prefix=""): selected_fields = [] for field in df.schema.fields: new_col_name = f"{parent_prefix}_{field.name}" if parent_prefix else field.name if isinstance(field.dataType, StructType): nested_df = df.select(col(field.name).alias(new_col_name)).select(f"{new_col_name}.*") selected_fields.extend(flatten_array_structs_with_prefix(nested_df, new_col_name).columns) elif isinstance(field.dataType, ArrayType) and isinstance(field.dataType.elementType, StructType): # 用transform函数处理数组内的Struct,生成新的数组列,每个元素是扁平化后的Struct transformed_array = transform( col(field.name), lambda elem: struct(*[elem[sub_field].alias(f"{new_col_name}_{sub_field}") for sub_field in field.dataType.elementType.fieldNames()]) ) selected_fields.append(transformed_array.alias(new_col_name)) else: selected_fields.append(col(field.name).alias(new_col_name)) return df.select(selected_fields)
测试后你会看到recent_orders还是数组类型,但里面的元素变成了recent_orders_order_id和recent_orders_total这样的字段。
3. 自定义调整
- 如果想要用其他分隔符(比如驼峰式
userProfileFullName),只需要修改new_col_name的生成逻辑,比如:new_col_name = parent_prefix + field.name.capitalize() if parent_prefix else field.name - 如果有多层嵌套,递归函数会自动处理,不需要额外修改。
内容的提问来源于stack exchange,提问作者Jas
相关产品推荐
相关产品推荐

