base.py 2.3 KB

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