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

如何在PySpark DataFrame架构中将LongType转换为ByteType?

修改PySpark DataFrame Schema中的LongType为ByteType

原Schema信息

1. df.printSchema()输出:

root
 |-- C_0_0: double (nullable = true)
 |-- C_0_1: array (nullable = true)
 |    |-- element: struct (containsNull = true)
 |    |    |-- C_2_0: array (nullable = true)
 |    |    |    |-- element: decimal(10,6) (containsNull = true)
 |    |    |-- C_2_1: array (nullable = true)
 |    |    |    |-- element: long (containsNull = true)
 |-- C_0_2: array (nullable = true)
 |    |-- element: struct (containsNull = true)
 |    |    |-- C_2_0: struct (nullable = true)
 |    |    |    |-- C_3_0: decimal(10,6) (nullable = true)
 |    |    |    |-- C_3_1: double (nullable = true)
 |    |    |-- C_2_1: struct (nullable = true)
 |    |    |    |-- C_3_0: timestamp (nullable = true)
 |    |    |    |-- C_3_1: double (nullable = true)

2. df.schema输出:

Schema:  StructType([StructField('C_0_0', DoubleType(), True), StructField('C_0_1', ArrayType(StructType([StructField('C_2_0', ArrayType(DecimalType(10,6), True), True), StructField('C_2_1', ArrayType(LongType(), True), True)]), True), True), StructField('C_0_2', ArrayType(StructType([StructField('C_2_0', StructType([StructField('C_3_0', DecimalType(10,6), True), StructField('C_3_1', DoubleType(), True)]), True), StructField('C_2_1', StructType([StructField('C_3_0', TimestampType(), True), StructField('C_3_1', DoubleType(), True)]), True)]), True), True)])

3. 期望的基础数据类型:

Basic Datatypes: ['DOUBLE', 'DECIMAL(10,6)', 'SMALLINT', 'DECIMAL(10,6)', 'DOUBLE', 'TIMESTAMP', 'DOUBLE']
注:PySpark不支持SMALLINT,需替换为ByteType。

4. Schema的字典形式:

{'fields': [
    {'metadata': {}, 'name': 'C_0_0', 'nullable': True, 'type': 'double'}, 
    {'metadata': {}, 'name': 'C_0_1', 'nullable': True, 'type': {
        'containsNull': True, 'elementType': {
            'fields': [
                {'metadata': {}, 'name': 'C_2_0', 'nullable': True, 'type': {'containsNull': True, 'elementType': 'double', 'type': 'array'}}, 
                {'metadata': {}, 'name': 'C_2_1', 'nullable': True, 'type': {'containsNull': True, 'elementType': 'long', 'type': 'array'}}
            ], 'type': 'struct'}, 'type': 'array'}}, 
    {'metadata': {}, 'name': 'C_0_2', 'nullable': True, 'type': {
        'containsNull': True, 'elementType': {
            'fields': [
                {'metadata': {}, 'name': 'C_2_0', 'nullable': True, 'type': {
                    'fields': [
                        {'metadata': {}, 'name': 'C_3_0', 'nullable': True, 'type': 'double'}, 
                        {'metadata': {}, 'name': 'C_3_1', 'nullable': True, 'type': 'double'}
                    ], 'type': 'struct'}}, 
                {'metadata': {}, 'name': 'C_2_1', 'nullable': True, 'type': {
                    'fields': [
                        {'metadata': {}, 'name': 'C_3_0', 'nullable': True, 'type': 'string'}, 
                        {'metadata': {}, 'name': 'C_3_1', 'nullable': True, 'type': 'double'}
                    ], 'type': 'struct'}}
            ], 'type': 'struct'}, 'type': 'array'}}
], 'type': 'struct'}

需求

将Schema中的LongType()(对应字典中的'long')修改为ByteType(),同时处理SMALLINT到ByteType的映射。

现有代码

def checkforbase(self, val):
    if isinstance(val,dict) or val=='array' or val == 'struct':
        return False
    else:
        return True

def modifySchema(self, schema):
    global cnt
    
    if 'type' in schema:
        if self.checkforbase(schema['type']):
            if basic_datatypes[cnt] == "SMALLINT":
                basic_datatypes[cnt] = "LONG"
            schema['type'] = basic_datatypes[cnt].lower()
            cnt = cnt + 1
        elif isinstance(schema['type'], dict):
            # print("Recursive again: ", d['type'])
            self.modifySchema(schema['type'])
    if 'elementType' in schema:
        if self.checkforbase(schema['elementType']):
            if basic_datatypes[cnt] == "SMALLINT":
                basic_datatypes[cnt] = "LONG"
            schema['elementType'] = basic_datatypes[cnt].lower()
            cnt = cnt + 1
        elif isinstance(schema['elementType'], dict):
            self.modifySchema(schema['elementType'])
    if 'fields' in schema:
        full_field = schema['fields']
        for i in range(len(full_field)):
            self.modifySchema(full_field[i])

解决方案

方案1:直接修改字典Schema

递归遍历字典结构,替换所有'long'类型为'byte',同时处理SMALLINT映射:

def update_schema_dict(schema_dict):
    if isinstance(schema_dict, dict):
        # 处理当前层级的type字段
        if 'type' in schema_dict:
            if schema_dict['type'] == 'long':
                schema_dict['type'] = 'byte'
            elif isinstance(schema_dict['type'], dict):
                update_schema_dict(schema_dict['type'])
        # 处理数组元素类型
        if 'elementType' in schema_dict:
            if schema_dict['elementType'] == 'long':
                schema_dict['elementType'] = 'byte'
            elif isinstance(schema_dict['elementType'], dict):
                update_schema_dict(schema_dict['elementType'])
        # 处理结构体字段数组
        if 'fields' in schema_dict:
            for field in schema_dict['fields']:
                update_schema_dict(field)
    return schema_dict

# 使用示例
updated_schema_dict = update_schema_dict(原始Schema字典)
# 转换回PySpark Schema对象
from pyspark.sql.types import StructType
updated_schema = StructType.fromJson(updated_schema_dict)

方案2:直接操作PySpark Schema对象

递归遍历Schema对象,替换LongType为ByteType:

from pyspark.sql.types import (
    StructType, StructField, ArrayType, ByteType,
    DoubleType, DecimalType, TimestampType
)

def update_schema(schema):
    if isinstance(schema, StructType):
        return StructType([update_schema(field) for field in schema.fields])
    elif isinstance(schema, StructField):
        return StructField(
            schema.name,
            update_schema(schema.dataType),
            schema.nullable,
            schema.metadata
        )
    elif isinstance(schema, ArrayType):
        element_type = update_schema(schema.elementType)
        return ArrayType(element_type, schema.containsNull)
    # 替换LongType为ByteType
    elif str(schema) == 'LongType':
        return ByteType()
    # 保留其他类型不变
    else:
        return schema

# 使用示例
updated_schema = update_schema(df.schema)
# 用新Schema转换DataFrame或读取数据

方案3:优化现有递归代码

调整原有逻辑,直接替换long为byte,同时处理SMALLINT映射:

class SchemaModifier:
    def __init__(self, basic_datatypes):
        self.basic_datatypes = basic_datatypes
        self.cnt = 0

    def checkforbase(self, val):
        return not (isinstance(val, dict) or val in ['array', 'struct'])

    def modifySchema(self, schema):
        if 'type' in schema:
            if self.checkforbase(schema['type']):
                target_type = self.basic_datatypes[self.cnt]
                # 处理SMALLINT映射为ByteType
                if target_type == "SMALLINT":
                    schema['type'] = 'byte'
                else:
                    schema['type'] = target_type.lower()
                self.cnt += 1
            elif isinstance(schema['type'], dict):
                self.modifySchema(schema['type'])
        if 'elementType' in schema:
            if self.checkforbase(schema['elementType']):
                target_type = self.basic_datatypes[self.cnt]
                if target_type == "SMALLINT":
                    schema['elementType'] = 'byte'
                elif schema['elementType'] == 'long':
                    schema['elementType'] = 'byte'
                else:
                    schema['elementType'] = target_type.lower()
                self.cnt += 1
            elif isinstance(schema['elementType'], dict):
                self.modifySchema(schema['elementType'])
        if 'fields' in schema:
            for field in schema['fields']:
                self.modifySchema(field)

# 使用示例
basic_datatypes = ['DOUBLE', 'DECIMAL(10,6)', 'SMALLINT', 'DECIMAL(10,6)', 'DOUBLE', 'TIMESTAMP', 'DOUBLE']
modifier = SchemaModifier(basic_datatypes)
updated_schema_dict = modifier.modifySchema(原始Schema字典)
updated_schema = StructType.fromJson(updated_schema_dict)

内容的提问来源于stack exchange,提问作者Abhik NASKAR

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.31 11:35:21