使用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,完全规避多生成器冲突问题。
预期输出
执行后会生成两个数组的笛卡尔积结果,示例如下:
| Child0 | Child1_Child2_name | Child1_Child2_Child21 | Child1_Child3_name | Child1_Child3_Child31 |
|---|---|---|---|---|
| 5678 | HRA | test 1 | LRA | test 3 |
| 5678 | HRA | test 1 | LRA | test 4 |
| 5678 | HRA | test 2 | LRA | test 3 |
| 5678 | HRA | test 2 | LRA | test 4 |
内容的提问来源于stack exchange,提问作者sameer A
相关产品推荐
相关产品推荐

