You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

使用PySpark展平含同级数组的XML时遇Multi Generator异常求助

PySpark展平含同级数组XML时的Multi Generator问题解决

错误信息

AnalysisException: [UNSUPPORTED_GENERATOR.MULTI_GENERATOR] The generator is not supported: only one generator allowed per SELECT clause but found 2: "generatorouter(explode(Child1.Child2.Child21))", "generatorouter(explode(Child1.Child3.Child31))".

XML结构

<Body>
<Child0>5678</Child0>
<Child1>
    <Child2 name="HRA">
        <Child21>test 1</Child21>
        <Child21>test 2</Child21>
    </Child2>
    <Child3 name="LRA">
        <Child31>test 3</Child31>
        <Child31>test 4</Child31>
    </Child3>
</Child1>
</Body>

现有实现代码

from pyspark.sql import SparkSession
from pyspark.sql.types import StructType, StructField, ArrayType
from pyspark.sql.functions import explode_outer

def flatten(df):
    f_df = df
    select_expr = _explodeArrays(element=f_df.schema)
    # While there is at least one Array, explode.
    while "ArrayType(" in f"{f_df.schema}": 
        f_df=f_df.selectExpr(select_expr)
        select_expr = _explodeArrays(element=f_df.schema)

    # Flatten the structure
    select_expr = flattenExpr(f_df.schema)
    f_df = f_df.selectExpr(select_expr)
    return f_df    

def _explodeArrays(element, root=None):
    el_type = type(element)
    expr = []
    try:
        _path = f"{root+'.' if root else ''}{element.name}"
    except AttributeError:
        _path = ""

    if el_type == StructType:
        for t in element:
            res = _explodeArrays(t, root)
            expr.extend(res)
    elif el_type == StructField and type(element.dataType) == ArrayType:
        expr.append(f"explode_outer({_path}) as {_path.replace('.','_')}")
    elif el_type == StructField and type(element.dataType) == StructType:
        expr.extend(_explodeArrays(element.dataType, _path))
    else:   
        expr.append(f"{_path}  as {_path.replace('.','_')}")

    return expr

def flattenExpr(element, root=None):
    expr = []
    el_type = type(element)
    try:
        _path = f"{root+'.' if root else ''}{element.name}"
    except AttributeError:
        _path = ""
    if el_type == StructType:
        for t in element:
            expr.extend(flattenExpr(t, root))
    elif el_type == StructField and type(element.dataType) == StructType:
        expr.extend(flattenExpr(element.dataType, _path))
    elif el_type == StructField and type(element.dataType) == ArrayType:
        # You should use flattenArrays to be sure this will not happen
        expr.extend(flattenExpr(element.dataType.elementType, f"{_path}[0]"))
    else:
        expr.append(f"{_path} as {_path.replace('.','_')}")
    return expr


spark = SparkSession.builder.getOrCreate()

path = 'Files/Test9.xml'
df = spark.read.format('xml').options(rowTag='Body', ignoreNamespace='true').load(path)

display('******* Initial Data Frame of XML file ********')
display(df)

display('******* Initial Schema of XML file ********')
df.printSchema()

f_df = flatten(df)

display('******* Flatten Schema of XML file ********')
f_df.printSchema()

display('******* Flatten  Data Frame of XML file ********')
display(f_df)

问题原因

原有代码在一次SELECT语句中尝试同时explode两个同级数组(Child21和Child31),但PySpark不允许单个SELECT子句中使用多个生成器函数(explode属于生成器)。

解决方案

修改展平逻辑,每次仅处理一个数组字段,循环执行直到所有数组都被展开,避免单次SELECT中出现多个生成器。

修改后的完整代码

from pyspark.sql import SparkSession
from pyspark.sql.types import StructType, StructField, ArrayType
from pyspark.sql.functions import explode_outer

def flatten(df):
    f_df = df
    # 循环处理每个数组字段,每次仅处理一个
    while True:
        array_field_path = find_first_array_field(f_df.schema)
        if not array_field_path:
            break
        # 生成SELECT表达式:保留所有非数组字段,仅explode当前找到的数组
        select_expr = []
        # 添加非数组字段
        for field in get_all_non_array_fields(f_df.schema):
            select_expr.append(f"{field} as {field.replace('.','_')}")
        # 添加当前数组的explode语句
        alias_name = array_field_path.replace('.', '_')
        select_expr.append(f"explode_outer({array_field_path}) as {alias_name}")
        # 执行SELECT
        f_df = f_df.selectExpr(select_expr)
    
    # 最后展平所有嵌套结构
    select_expr = flattenExpr(f_df.schema)
    f_df = f_df.selectExpr(select_expr)
    return f_df

def find_first_array_field(schema, root_path=""):
    """递归查找第一个数组类型字段的完整路径"""
    for field in schema.fields:
        current_path = f"{root_path}.{field.name}" if root_path else field.name
        if isinstance(field.dataType, ArrayType):
            return current_path
        elif isinstance(field.dataType, StructType):
            nested_path = find_first_array_field(field.dataType, current_path)
            if nested_path:
                return nested_path
    return None

def get_all_non_array_fields(schema, root_path=""):
    """获取所有非数组类型字段的完整路径"""
    fields = []
    for field in schema.fields:
        current_path = f"{root_path}.{field.name}" if root_path else field.name
        if isinstance(field.dataType, ArrayType):
            continue
        elif isinstance(field.dataType, StructType):
            fields.extend(get_all_non_array_fields(field.dataType, current_path))
        else:
            fields.append(current_path)
    return fields

def flattenExpr(element, root=None):
    expr = []
    el_type = type(element)
    try:
        _path = f"{root+'.' if root else ''}{element.name}"
    except AttributeError:
        _path = ""
    if el_type == StructType:
        for t in element:
            expr.extend(flattenExpr(t, root))
    elif el_type == StructField and type(element.dataType) == StructType:
        expr.extend(flattenExpr(element.dataType, _path))
    elif el_type == StructField and type(element.dataType) == ArrayType:
        expr.extend(flattenExpr(element.dataType.elementType, f"{_path}[0]"))
    else:
        expr.append(f"{_path} as {_path.replace('.','_')}")
    return expr


spark = SparkSession.builder.getOrCreate()

path = 'Files/Test9.xml'
df = spark.read.format('xml').options(rowTag='Body', ignoreNamespace='true').load(path)

display('******* Initial Data Frame of XML file ********')
display(df)

display('******* Initial Schema of XML file ********')
df.printSchema()

f_df = flatten(df)

display('******* Flatten Schema of XML file ********')
f_df.printSchema()

display('******* Flatten  Data Frame of XML file ********')
display(f_df)

代码说明

  • find_first_array_field:递归查找第一个数组字段的完整路径,确保每次仅处理一个数组。
  • get_all_non_array_fields:收集所有非数组字段,在explode时保留这些字段的原始值。
  • 优化后的flatten函数:循环处理每个数组,每次仅对一个数组执行explode_outer,完全规避多生成器冲突问题。

预期输出

执行后会生成两个数组的笛卡尔积结果,示例如下:

Child0Child1_Child2_nameChild1_Child2_Child21Child1_Child3_nameChild1_Child3_Child31
5678HRAtest 1LRAtest 3
5678HRAtest 1LRAtest 4
5678HRAtest 2LRAtest 3
5678HRAtest 2LRAtest 4

内容的提问来源于stack exchange,提问作者sameer A

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.06.25 06:50:10