如何在Pydantic中实现FastAPI端点对scipy lil_matrix的支持?
让Pydantic支持Scipy lil_matrix并适配FastAPI的最小实现步骤
要在FastAPI中直接使用Scipy的lil_matrix类型,无需转稠密矩阵浪费资源,你需要实现以下核心内容:
1. 定义自定义Pydantic兼容的LilMatrix类型
通过实现Pydantic v2的核心schema接口,告诉框架如何验证和序列化lil_matrix:
from scipy.sparse import lil_matrix from pydantic import GetCoreSchemaHandler, ValidationInfo from pydantic_core import core_schema class LilMatrix(lil_matrix): @classmethod def __get_pydantic_core_schema__(cls, source_type, handler: GetCoreSchemaHandler) -> core_schema.CoreSchema: # 验证逻辑:处理lil_matrix实例或COO格式字典 def validate_input(value, info: ValidationInfo): if isinstance(value, lil_matrix): return value # 从COO字典反序列化为lil_matrix elif isinstance(value, dict) and all(key in value for key in ("data", "row", "col", "shape")): return lil_matrix( (value["data"], (value["row"], value["col"])), shape=value["shape"] ) raise ValueError("输入必须是lil_matrix实例或包含data、row、col、shape的COO格式字典") # 序列化逻辑:将lil_matrix转为COO格式字典(可复用你写的sparse_to_coo_dict逻辑) def serialize_output(value): coo_matrix = value.tocoo() return { "data": coo_matrix.data.tolist(), "row": coo_matrix.row.tolist(), "col": coo_matrix.col.tolist(), "shape": coo_matrix.shape } # 组合核心Schema:支持两种输入类型,指定序列化规则 return core_schema.no_info_after_validator_function( validate_input, core_schema.union_schema([ core_schema.is_instance_schema(lil_matrix), core_schema.dict_schema() ]), serialization=core_schema.plain_serializer_function_ser_schema(serialize_output) )
2. 在Pydantic模型中使用自定义类型
直接把LilMatrix作为字段类型,FastAPI会自动处理请求解析和响应序列化:
from pydantic import BaseModel class SparseMatrixRequest(BaseModel): matrix: LilMatrix scalar_multiplier: float class SparseMatrixResponse(BaseModel): result: LilMatrix
3. 编写FastAPI端点
在接口中直接使用自定义模型,无需额外转换即可调用Scipy的线性代数方法:
from fastapi import FastAPI app = FastAPI() @app.post("/sparse-multiply") def multiply_sparse_matrix(request: SparseMatrixRequest): # 直接对lil_matrix执行Scipy运算 result = request.matrix * request.scalar_multiplier return SparseMatrixResponse(result=result)
关键说明
- 验证逻辑同时支持两种输入:后端内部传递的
lil_matrix实例,以及前端/客户端发送的COO格式字典 - 序列化逻辑完全基于稀疏格式的原生数据传输,避免了稠密矩阵转换带来的资源浪费
- 你之前实现的
sparse_to_coo_dict可以直接替换到serialize_output函数中,保持逻辑一致性
测试示例请求
发送POST请求到/sparse-multiply,请求体格式如下:
{ "matrix": { "data": [1.5, 3.2, 2.7], "row": [0, 1, 2], "col": [2, 0, 1], "shape": [3, 3] }, "scalar_multiplier": 2.0 }
内容的提问来源于stack exchange,提问作者Chechy Levas
相关产品推荐
相关产品推荐

