Pydantic:让 Python 数据有“合同”

一句话:Pydantic 是 Python 里用来做数据校验、类型转换和配置管理的库。

你写 Python 时大概率遇到过这些问题:

def train(model_name, batch_size, learning_rate):
    ...

调用的人可能传进来:

train("bert", "32", "0.001")
train(None, -1, 1e-5)
train({"name": "bert"}, 32.5, 0.01)

函数内部不知道到底拿到了什么。
类型注解只能“看起来很美”,运行时并不一定拦得住。

Pydantic 的作用就是:把数据变成有约束、有类型、有默认值、可校验、可序列化的对象。


1. 最小例子

from pydantic import BaseModel, Field
 
 
class TrainConfig(BaseModel):
    model_name: str
    batch_size: int = Field(default=32, gt=0)
    learning_rate: float = Field(default=1e-4, gt=0)
    use_amp: bool = True
 
 
config = TrainConfig(
    model_name="bert",
    batch_size="64",
    learning_rate="0.001"
)
 
print(config)

输出大概像这样:

TrainConfig(model_name='bert', batch_size=64, learning_rate=0.001, use_amp=True)

注意两点:

  1. "64" 被自动转成了 int
  2. "0.001" 被自动转成了 float

如果传入非法值:

TrainConfig(model_name="bert", batch_size=-1)

会直接报错:

ValidationError: 1 validation error for TrainConfig
batch_size
  Input should be greater than 0

这就是 Pydantic 的核心价值:
数据进入系统时先签合同,不合格直接拒绝。


2. 为什么 AI/算法项目特别需要它?

算法代码通常不是死在模型里,而是死在数据和配置里。

比如:

  • 配置文件字段写错
  • 超参数类型不对
  • 数据集路径为空
  • API 返回 JSON 字段缺失
  • LLM 输出不是合法 JSON
  • 推理服务请求参数不合法
  • 训练参数混入脏数据

这些问题用 if/else 判断会越写越臭。
Pydantic 更适合做一层统一的数据边界。

你可以把它理解成:

Python 里的数据结构 + 类型校验 + JSON Schema + 序列化。


3. 常见用法:定义请求参数

做模型服务时,经常要定义输入结构:

from pydantic import BaseModel, Field
 
 
class PredictRequest(BaseModel):
    text: str = Field(min_length=1, max_length=512)
    top_k: int = Field(default=5, ge=1, le=100)
    temperature: float = Field(default=0.7, ge=0.0, le=2.0)

用户请求:

req = PredictRequest(text="介绍一下 Pydantic", top_k="10")

Pydantic 会自动校验并转换:

print(req.top_k)
print(type(req.top_k))

结果:

10
<class 'int'>

如果用户传:

PredictRequest(text="", top_k=1000)

直接失败:

ValidationError

这比你在接口里手写一堆校验逻辑清爽很多。


4. 自定义校验器

有时候不是简单类型检查,而是业务规则。

比如:模型名不能为空,并且必须以 .pt.bin 结尾。

from pydantic import BaseModel, field_validator
 
 
class ModelConfig(BaseModel):
    model_path: str
 
    @field_validator("model_path")
    @classmethod
    def check_model_path(cls, v: str) -> str:
        if not v.strip():
            raise ValueError("model_path 不能为空")
 
        if not v.endswith((".pt", ".bin")):
            raise ValueError("model_path 必须是 .pt 或 .bin 文件")
 
        return v

测试:

ModelConfig(model_path="weights/bert.pt")  # OK
ModelConfig(model_path="weights/bert.onnx")  # 报错

注意:field_validator 里要返回最终字段值。


5. 跨字段校验

有时一个字段依赖另一个字段。

比如:如果开启分布式训练,必须指定 world_size

from pydantic import BaseModel, model_validator
 
 
class DistributedConfig(BaseModel):
    use_ddp: bool = False
    world_size: int | None = None
 
    @model_validator(mode="after")
    def check_world_size(self):
        if self.use_ddp and self.world_size is None:
            raise ValueError("开启 use_ddp 时必须设置 world_size")
        return self

这样就可以保证配置整体合理,而不是只校验单个字段。


6. JSON 序列化与反序列化

Pydantic 模型天然适合和 JSON 打交道。

class TrainConfig(BaseModel):
    model_name: str
    batch_size: int = 32
    learning_rate: float = 1e-4

对象转 dict:

config = TrainConfig(model_name="bert")
 
print(config.model_dump())

输出:

{
    "model_name": "bert",
    "batch_size": 32,
    "learning_rate": 0.0001
}

对象转 JSON:

print(config.model_dump_json())

JSON 转对象:

json_str = """
{
    "model_name": "bert",
    "batch_size": 64,
    "learning_rate": 0.001
}
"""
 
config = TrainConfig.model_validate_json(json_str)
print(config.batch_size)

输出:

64

这在读取配置文件、调用 API、解析 LLM 输出时特别有用。


7. 生成 JSON Schema

Pydantic 可以根据模型自动生成 JSON Schema。

print(TrainConfig.model_json_schema())

这在以下场景很有用:

  • OpenAPI 接口文档
  • LLM function calling
  • structured output
  • 表单生成
  • 配置校验
  • 自动化测试

例如让 LLM 按固定结构输出时,可以把 Pydantic schema 给模型:

schema = TrainConfig.model_json_schema()

然后要求模型输出符合该 schema 的 JSON。


8. 在 LLM 应用里特别好用

现在做 LLM 应用,最怕模型“胡说八道”地返回结构。

你希望模型返回:

{
  "answer": "Pydantic 是数据校验库",
  "confidence": 0.92,
  "reasoning": "它基于类型注解进行运行时校验"
}

可以定义:

from pydantic import BaseModel, Field
 
 
class LLMAnswer(BaseModel):
    answer: str
    confidence: float = Field(ge=0, le=1)
    reasoning: str | None = None

拿到模型输出后:

raw_output = """
{
    "answer": "Pydantic 是数据校验库",
    "confidence": 0.92,
    "reasoning": "它基于类型注解进行运行时校验"
}
"""
 
result = LLMAnswer.model_validate_json(raw_output)
 
print(result.answer)
print(result.confidence)

如果模型返回:

{
  "answer": "不知道",
  "confidence": 2.5
}

Pydantic 会直接告诉你:confidence 不合法。

所以 Pydantic 很适合做:

  • LLM 结构化输出校验
  • Agent 工具调用参数校验
  • RAG 检索结果结构化
  • Prompt 配置管理
  • 模型 API 输入输出约束

9. 配置管理:pydantic-settings

训练、推理、服务部署经常需要读取环境变量。

Pydantic v2 里,配置管理通常用独立包:

pip install pydantic-settings

示例:

from pydantic_settings import BaseSettings
 
 
class Settings(BaseSettings):
    app_name: str = "ml-service"
    model_path: str
    batch_size: int = 32
    debug: bool = False
 
    class Config:
        env_file = ".env"

如果 .env 文件里有:

MODEL_PATH=/data/models/bert.pt
BATCH_SIZE=64
DEBUG=true

那么:

settings = Settings()
 
print(settings.model_path)
print(settings.batch_size)
print(settings.debug)

输出:

/data/models/bert.pt
64
True

它会自动读取环境变量并做类型转换。

这比手写 os.getenv 干净太多。


10. 和 FastAPI 是绝配

FastAPI 的请求体校验就是基于 Pydantic。

from fastapi import FastAPI
from pydantic import BaseModel, Field
 
app = FastAPI()
 
 
class PredictRequest(BaseModel):
    text: str = Field(min_length=1)
    top_k: int = Field(default=5, ge=1, le=100)
 
 
@app.post("/predict")
def predict(req: PredictRequest):
    return {
        "text": req.text,
        "top_k": req.top_k
    }

用户传错参数时,FastAPI 会自动返回 422 错误。
你不用手写参数校验。

如果你做模型推理服务、数据标注平台、LLM API 网关,这套组合非常常见。


11. 常见坑

坑 1:默认值不要用可变对象

错误:

class Config(BaseModel):
    tags: list[str] = []

更稳:

from pydantic import Field
 
 
class Config(BaseModel):
    tags: list[str] = Field(default_factory=list)

坑 2:validator 必须返回值

错误:

@field_validator("name")
@classmethod
def check_name(cls, v):
    if not v:
        raise ValueError("name 不能为空")

应该:

@field_validator("name")
@classmethod
def check_name(cls, v):
    if not v:
        raise ValueError("name 不能为空")
    return v

坑 3:不要把所有东西都塞进 Pydantic

Pydantic 适合校验结构化数据,不适合直接校验超大文件、超大张量、超大数据集。

比如你有 10GB 图像数据,不要指望把原始图片塞进 Pydantic 模型里做全量校验。

更合理的是:

  • 元数据用 Pydantic
  • 文件路径用 Pydantic
  • 参数配置用 Pydantic
  • 大张量用专门的数据加载器处理

坑 4:输入模型和输出模型最好分开

很多新手喜欢一个模型从头用到尾。

更清晰的做法是分开:

class UserCreate(BaseModel):
    username: str
    password: str
 
 
class UserRead(BaseModel):
    id: int
    username: str

输入模型负责校验用户提交。
输出模型负责返回安全字段。

避免把密码、token、内部状态泄露出去。


12. 一个工程化小模板

在 AI 项目里,我通常会这样组织配置:

from pydantic import BaseModel, Field
 
 
class DataConfig(BaseModel):
    train_path: str
    val_path: str
    batch_size: int = Field(default=32, gt=0)
    num_workers: int = Field(default=4, ge=0)
 
 
class ModelConfig(BaseModel):
    name: str
    hidden_size: int = Field(default=768, gt=0)
    dropout: float = Field(default=0.1, ge=0.0, le=1.0)
 
 
class TrainConfig(BaseModel):
    learning_rate: float = Field(default=1e-4, gt=0)
    epochs: int = Field(default=10, gt=0)
    use_amp: bool = True
    seed: int = 42
 
    data: DataConfig
    model: ModelConfig

使用时:

config = TrainConfig(
    data={
        "train_path": "/data/train.jsonl",
        "val_path": "/data/val.jsonl",
        "batch_size": "64"
    },
    model={
        "name": "bert",
        "hidden_size": 768,
        "dropout": 0.2
    },
    learning_rate="0.0001"
)
 
print(config.data.batch_size)
print(config.model.dropout)

嵌套结构也能自动校验。
这非常适合复杂训练配置。


13. Pydantic 的本质

很多人把 Pydantic 理解成“类型检查工具”,这不够准确。

它更像一个运行时数据契约系统:

原始数据

Pydantic 校验

合法对象

业务逻辑

它保证进入业务逻辑的数据是可信的。

对于 AI 工程来说,这特别重要:

  • 训练配置可信
  • 推理请求可信
  • 模型输出可信
  • API 数据可信
  • 实验参数可信

很多 bug 不是算法问题,而是数据没守住边界。
Pydantic 就是帮你守边界。


14. 什么时候该用 Pydantic?

推荐使用:

  • API 请求/响应结构定义
  • 训练/推理配置
  • 环境变量读取
  • LLM 结构化输出
  • Agent 工具参数
  • 数据集元信息
  • 实验参数记录
  • 模型服务输入校验

不推荐使用:

  • 超大规模数据逐条校验
  • 高性能数值计算核心路径
  • 替代 NumPy/Pandas 做数据分析
  • 替代 PyTorch/TensorFlow 做张量处理

15. 总结

Pydantic 不是让模型变强的库,但它是让 AI 工程更稳的库。

它的核心价值可以概括为三句话:

1. 用类型注解定义数据结构
2. 在运行时自动校验和转换数据
3. 方便地序列化、反序列化和生成 Schema

如果你正在做:

  • 模型服务
  • LLM 应用
  • RAG 系统
  • Agent 工具调用
  • 训练配置管理
  • FastAPI 接口
  • 数据管道入口校验

那么 Pydantic 基本值得成为默认工具。

最简单的上手方式:

pip install pydantic

然后从一个配置类开始:

from pydantic import BaseModel, Field
 
 
class TrainConfig(BaseModel):
    model_name: str
    batch_size: int = Field(default=32, gt=0)
    learning_rate: float = Field(default=1e-4, gt=0)

当你第一次因为 Pydantic 提前拦下脏数据时,你会明白它的意义。