如何在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

