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

如何在Pandera SchemaModel继承中自定义必填与可选字段?

问题描述

我定义了InputSchema和OutputSchema两个Pandera SchemaModel,其中OutputSchema继承自InputSchema。需求如下:

  • InputSchema中的字段保持必填状态
  • OutputSchema仅保留部分继承字段为必填,其余自动转为可选

尝试参考Pydantic的解决方案,在父类中定义__init_subclass__方法,但运行时出现KeyError: 'reporting_date'错误,测试代码如下:

import pandas as pd
import pandera as pa
from typing import Optional
from pandera.typing import DataFrame, Series

class InputSchema(pa.SchemaModel):
    reporting_date: Series[pa.DateTime] = pa.Field(coerce=True)

    def __init_subclass__(cls, optional_fields=None, **kwargs):
        super().__init_subclass__(**kwargs)
        if optional_fields:
            for field in optional_fields:
                cls.__fields__[field].outer_type_ = Optional
                cls.__fields__[field].required = False    

class OutputSchema(InputSchema, optional_fields=['reporting_date']):
    test: Series[str] = pa.Field()

@pa.check_types
def func(inputs: DataFrame[InputSchema]) -> DataFrame[OutputSchema]:
    inputs = inputs.drop(columns=['reporting_date'])
    inputs['test'] = 'a'
    return inputs

data = pd.DataFrame({'reporting_date': ['2023-01-11', '2023-01-12']})
func(data)

报错信息:

--> 18 class OutputSchema(InputSchema, optional_fields=['reporting_date']):
KeyError: 'reporting_date'

期望实现的效果是:可以在子类中指定继承字段的必填项,其余自动转为可选,示例写法如下:

class InputSchema(pa.SchemaModel):
    reporting_date: Series[pa.DateTime] = pa.Field(coerce=True)
    other_field: Series[str] = pa.Field()

class OutputSchema(InputSchema, required=['reporting_date']):
    test: Series[str] = pa.Field()

最终OutputSchema中reporting_date和test为必填字段,other_field为可选字段。

解决方案

Pandera的SchemaModel内部机制和Pydantic不同,直接操作__fields__会因为字段初始化时机问题报错。可以通过以下两种方式实现需求:

方法一:重定义父类字段并标记为可选

在子类中重新声明需要转为可选的父类字段,将其类型改为Optional,并设置pa.Field(required=False):

import pandas as pd
import pandera as pa
from typing import Optional
from pandera.typing import DataFrame, Series

class InputSchema(pa.SchemaModel):
    reporting_date: Series[pa.DateTime] = pa.Field(coerce=True)
    other_field: Series[str] = pa.Field()

class OutputSchema(InputSchema):
    # 重定义需要转为可选的字段
    other_field: Optional[Series[str]] = pa.Field(required=False)
    # 新增的必填字段
    test: Series[str] = pa.Field()

@pa.check_types
def func(inputs: DataFrame[InputSchema]) -> DataFrame[OutputSchema]:
    inputs = inputs.drop(columns=['other_field'])
    inputs['test'] = 'a'
    return inputs

data = pd.DataFrame({
    'reporting_date': ['2023-01-11', '2023-01-12'],
    'other_field': ['foo', 'bar']
})
result = func(data)
print(result)

这种方式直观,适合字段数量不多的场景。

方法二:自定义父类的字段处理逻辑

利用Pandera的SchemaModel的__init_subclass__结合__post_init__钩子,延迟到类初始化完成后修改字段必填性,实现你期望的required_fields参数语法:

import pandas as pd
import pandera as pa
from typing import Optional
from pandera.typing import DataFrame, Series

class BaseSchema(pa.SchemaModel):
    def __init_subclass__(cls, required_fields=None, **kwargs):
        super().__init_subclass__(**kwargs)
        cls._required_fields = required_fields or []

    @classmethod
    def __post_init__(cls):
        # 类初始化完成后,所有字段已加载完成,再修改必填性
        if hasattr(cls, '_required_fields'):
            # 获取所有父类的SchemaModel字段
            parent_fields = {}
            for base in cls.__bases__:
                if issubclass(base, pa.SchemaModel) and base != pa.SchemaModel:
                    parent_fields.update(base.__fields__)
            
            # 将不在required_fields中的父类字段转为可选
            for field_name, field in parent_fields.items():
                if field_name not in cls._required_fields:
                    # 更新字段注解为Optional类型
                    cls.__annotations__[field_name] = Optional[cls.__annotations__[field_name]]
                    # 修改字段的必填属性
                    cls.__fields__[field_name].required = False
                    # 更新字段的类型信息
                    cls.__fields__[field_name].type_ = Optional[field.type_]

class InputSchema(BaseSchema):
    reporting_date: Series[pa.DateTime] = pa.Field(coerce=True)
    other_field: Series[str] = pa.Field()

class OutputSchema(InputSchema, required_fields=['reporting_date']):
    test: Series[str] = pa.Field()

@pa.check_types
def func(inputs: DataFrame[InputSchema]) -> DataFrame[OutputSchema]:
    inputs = inputs.drop(columns=['other_field'])
    inputs['test'] = 'a'
    return inputs

data = pd.DataFrame({
    'reporting_date': ['2023-01-11', '2023-01-12'],
    'other_field': ['foo', 'bar']
})
result = func(data)
print(result)

这种方式自动将父类中不在required_fields列表里的字段转为可选,适合字段较多的场景。

原理说明

Pandera的SchemaModel在类定义时会先收集注解和字段,然后在__post_init__阶段完成字段的最终初始化。直接在__init_subclass__中操作__fields__会报错,因为此时子类的__fields__还未完全从父类继承并初始化。通过__post_init__钩子,可以确保所有字段都已加载完成,再进行修改。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.03 23:55:32