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

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.13 01:35:09