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

如何在Pydantic BaseModel中为第三方类实现泛型子类型注解支持

问题描述

已实现一个Pydantic注解包装器,可将JSON列表解析为NumPy数组,但无法支持指定dtype的校验需求——例如希望使用NumpyWrapper[np.float64]这样的语法同时校验数组类型。尝试用泛型改造时出现错误:‘类型“ndarray[Any, dtype[Unknown]]”已被特化’。

解决方案

要实现支持指定dtype的NumPy数组注解包装器,需要正确结合泛型与Pydantic的核心Schema机制,同时在验证逻辑中加入dtype校验。以下是修改后的完整代码:

from typing import Annotated, Any, Generic, TypeVar
import numpy as np
import numpy.typing as npt
from pydantic import GetCoreSchemaHandler, GetJsonSchemaHandler
from pydantic.json_schema import JsonSchemaValue
from pydantic_core import core_schema

# 定义受Numpy dtype约束的类型变量
DTypeLike = TypeVar("DTypeLike", bound=npt.DTypeLike)

class _NumpyPydanticAnnotation(Generic[DTypeLike]):
    @classmethod
    def __get_pydantic_core_schema__(
        cls,
        source_type: Any,
        handler: GetCoreSchemaHandler,
    ) -> core_schema.CoreSchema:
        # 从泛型参数中提取目标dtype
        target_dtype = source_type.__args__[0] if hasattr(source_type, "__args__") else None

        def validate_from_list(value: list) -> np.ndarray:
            arr = np.array(value, dtype=target_dtype)
            # 校验实际生成的数组dtype是否匹配目标dtype(处理自动类型提升的情况)
            if not np.issubdtype(arr.dtype, target_dtype):
                raise ValueError(f"Array dtype {arr.dtype} does not match required dtype {target_dtype}")
            return arr

        def validate_numpy_array(value: np.ndarray) -> np.ndarray:
            if not np.issubdtype(value.dtype, target_dtype):
                raise ValueError(f"Array dtype {value.dtype} does not match required dtype {target_dtype}")
            return value

        from_list_schema = core_schema.chain_schema(
            [
                core_schema.list_schema(),
                core_schema.no_info_plain_validator_function(validate_from_list),
            ]
        )

        numpy_instance_schema = core_schema.chain_schema(
            [
                core_schema.is_instance_schema(np.ndarray),
                core_schema.no_info_plain_validator_function(validate_numpy_array),
            ]
        )

        return core_schema.json_or_python_schema(
            json_schema=from_list_schema,
            python_schema=core_schema.union_schema([numpy_instance_schema, from_list_schema]),
            serialization=core_schema.plain_serializer_function_ser_schema(lambda instance: instance.tolist()),
        )

    @classmethod
    def __get_pydantic_json_schema__(
        cls, _core_schema: core_schema.CoreSchema, handler: GetJsonSchemaHandler
    ) -> JsonSchemaValue:
        # 沿用列表的JSON Schema定义
        return handler(core_schema.list_schema())

# 定义泛型的Annotated包装类型
NumpyWrapper = Annotated[npt.NDArray[DTypeLike], _NumpyPydanticAnnotation[DTypeLike]]

关键说明

  • 泛型绑定:定义DTypeLike类型变量并绑定到npt.DTypeLike,确保传入的类型是合法的NumPy dtype。
  • 提取目标dtype:通过source_type.__args__从泛型参数中获取指定的目标dtype,作为校验依据。
  • 双重校验逻辑:
    • 从JSON列表解析时,直接用目标dtype创建数组,并校验实际生成的dtype是否匹配(避免自动类型提升导致的不符合预期)。
    • 对已有的NumPy数组实例,直接校验其dtype是否符合要求。
  • Schema兼容性:保持JSON Schema为列表类型,同时支持Python端的NumPy数组实例直接传入。

使用示例

from pydantic import BaseModel

class SomeDataModel(BaseModel):
    float_array: NumpyWrapper[np.float64]
    int_array: NumpyWrapper[np.int32]

# 合法场景
valid_data = {"float_array": [1.0, 2.5], "int_array": [1, 2]}
model = SomeDataModel(**valid_data)
print(model.float_array.dtype)  # 输出: float64
print(model.int_array.dtype)    # 输出: int32

# 非法场景:传入不兼容类型会触发错误
invalid_data = {"float_array": ["not a number"], "int_array": [1.5]}
try:
    model = SomeDataModel(**invalid_data)
except ValueError as e:
    print(e)  # 抛出dtype不匹配或数组创建失败的错误信息

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.30 20:21:08