| 12345678910111213141516171819202122232425262728293031323334353637383940414243444546474849505152535455565758596061626364 |
- # -*- coding: utf-8 -*-
- """Predictor 抽象基类。
- 按 ARCHITECTURE.md §6.2 Predictor 接口约束:
- - Predictor 不直接读写 Redis/PG。
- - 模型路径从配置读取,不拼接代码目录。
- - load() 与 predict() 分离,便于启动预热和测试 mock。
- - Predictor 之间不互相 import,通过 Pipeline 编排顺序传递结果。
- Phase 2 状态:
- 本文件仅定义接口契约,不强制现有 interface/predictor.py 中的 28 个
- Predictor 类立即继承。Phase 5 迁移 Predictor 时再逐一改为继承 BasePredictor。
- PredictorRegistry 不强制要求注册的实例是 BasePredictor 子类,
- 只要是可调用对象即可,以保证旧代码兼容。
- """
- from __future__ import absolute_import
- __all__ = ["BasePredictor"]
- class BasePredictor(object):
- """所有 Predictor 的抽象基类。
- 子类必须实现 :meth:`load` 和 :meth:`predict`。
- ``name`` 属性用于在 PredictorRegistry 中注册和查找,应与 registry key 一致。
- 示例::
- class CodeNamePredict(BasePredictor):
- name = "codeName"
- def load(self):
- # 加载模型权重、词表等
- ...
- def predict(self, ctx):
- # 执行推理,返回 dict 写回 ctx.model_results
- return {"code": ..., "name": ...}
- """
- #: Predictor 名称,对应 PredictorRegistry 的 key。
- name = ""
- def load(self):
- """加载模型权重、词表、配置等重型资源。
- 应与 :meth:`predict` 分离,便于:
- - 启动时预热(load 一次,predict 多次)
- - 测试时 mock(替换 load 为空实现)
- - 延迟加载(首次 predict 前才 load)
- 子类必须实现。本基类提供默认空实现,便于旧 Predictor 逐步过渡。
- """
- return None
- def predict(self, ctx):
- """执行推理。
- :param ctx: PipelineContext(Phase 2 仅占位,Phase 4/5 落地)。
- 旧 Predictor 不接收 ctx,可直接忽略此参数,保持原签名。
- :returns: dict,由调用方写回 ``ctx.model_results[self.name]``。
- 旧 Predictor 返回原结构即可,由 Stage 适配。
- """
- raise NotImplementedError("BasePredictor.predict must be overridden")
|