registry.py 12 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302303304305306307308309310311312313314315316
  1. # -*- coding: utf-8 -*-
  2. """Predictor 注册表(PredictorRegistry)。
  3. 替代 interface/predictor.py 中的全局 ``dict_predictor`` 字典。
  4. 设计目标(ARCHITECTURE.md §4.3、§6.2、Phase 2):
  5. 1. 线程安全:每个 predictor 有独立 RLock,与原 dict_predictor 行为一致。
  6. 2. 懒加载:首次 ``get(name)`` 才调用 factory 创建实例,避免启动时全量加载。
  7. 3. 显式注册:通过 ``register(name, factory)`` 注册,便于 Phase 5 逐个迁移 Predictor。
  8. 4. 兼容旧入口:``interface/predictor.py::getPredictor()`` 改为委托本注册表,
  9. 老 import 路径无需改动。
  10. 5. 不强制 BasePredictor 子类:只要是可调用 factory 返回实例即可,
  11. 保证现有 28 个 Predictor 类无需立即改造。
  12. 迁移路径:
  13. Phase 2(本文件):
  14. - 所有 28 个 Predictor 仍在 interface/predictor.py 中定义
  15. - register_defaults() 用 lambda 延迟 import 旧位置类
  16. - getPredictor() 委托 get_default_registry().get(name)
  17. Phase 5:
  18. - Predictor 类逐个迁移到 predictors/<name>.py
  19. - 每个 <name>.py 在 import 时调用 register() 覆盖默认 factory
  20. - interface/predictor.py 变成纯 re-export 兼容层,最终删除
  21. 用法示例::
  22. from BiddingKG.dl.predictors.registry import get_default_registry, register_defaults
  23. register_defaults() # 幂等,多次调用只注册一次
  24. predictor = get_default_registry().get("codeName")
  25. # Phase 5 自定义注册:
  26. from BiddingKG.dl.predictors.codename import CodeNamePredict
  27. registry = get_default_registry()
  28. registry.register("codeName", lambda: CodeNamePredict()) # 覆盖默认 factory
  29. """
  30. from __future__ import absolute_import
  31. from threading import RLock
  32. __all__ = [
  33. "PredictorRegistry",
  34. "get_default_registry",
  35. "register_defaults",
  36. "is_defaults_registered",
  37. ]
  38. #: 原 dict_predictor 中 28 个 predictor 的 key 列表(按原文件顺序)。
  39. #: 顺序仅用于文档展示,registry 内部用 dict 存储不依赖顺序。
  40. DEFAULT_PREDICTOR_NAMES = (
  41. "codeName",
  42. "prem",
  43. "epc",
  44. "roleRule",
  45. "roleRuleFinal",
  46. "tendereeRuleRecall",
  47. "form",
  48. "time",
  49. "punish",
  50. "product",
  51. "product_attrs",
  52. "channel",
  53. "deposit_payment_way",
  54. "total_unit_money",
  55. "industry",
  56. "rolegrade",
  57. "moneygrade",
  58. "district",
  59. "tableprem",
  60. "candidate",
  61. "websource_tenderee",
  62. "project_label",
  63. "industry_label",
  64. "pb_extract",
  65. "property_label",
  66. "approval",
  67. "bid_score",
  68. "entity_type_rule",
  69. )
  70. class PredictorRegistry(object):
  71. """线程安全的 Predictor 注册表。
  72. 与原 ``dict_predictor`` 的对应关系:
  73. - ``dict_predictor[name]["Lock"]`` -> 内部 ``self._locks[name]``
  74. - ``dict_predictor[name]["predictor"]`` -> 内部 ``self._instances[name]``
  75. - ``getPredictor(name)`` -> ``self.get(name)``
  76. factory 是零参数 callable,返回 Predictor 实例。
  77. factory 在 ``register()`` 时不被调用,只在首次 ``get()`` 时调用一次。
  78. """
  79. def __init__(self):
  80. self._factories = {}
  81. self._instances = {}
  82. self._locks = {}
  83. # 保护 _factories / _locks 字典本身的结构性变更(register/clear/reset)。
  84. # 单个 predictor 实例的创建由 self._locks[name] 保护。
  85. self._global_lock = RLock()
  86. # ------------------------------------------------------------------ register
  87. def register(self, name, factory):
  88. """注册或覆盖一个 predictor factory。
  89. :param name: predictor key,例如 "codeName"。
  90. :param factory: 零参数 callable,返回 predictor 实例。
  91. 重复 register 同名 key 会覆盖旧 factory,但已创建的实例不会立即失效,
  92. 需调用 :meth:`reset` 清除旧实例。
  93. 幂等性:重复 register 同一个 (name, factory) 不影响已创建实例。
  94. """
  95. if not callable(factory):
  96. raise TypeError("factory must be callable, got %r" % type(factory))
  97. with self._global_lock:
  98. self._factories[name] = factory
  99. if name not in self._locks:
  100. self._locks[name] = RLock()
  101. # --------------------------------------------------------------------- get
  102. def get(self, name):
  103. """获取 predictor 实例,首次调用时触发 factory 创建。
  104. :raises KeyError: name 未注册时抛出。
  105. 原实现抛 ``NameError("no this type of predictor")``,
  106. ``interface/predictor.py::getPredictor`` 会捕获并转换为 NameError
  107. 以保持兼容。
  108. """
  109. if name not in self._factories:
  110. raise KeyError("no predictor registered for: %s" % name)
  111. # 双重检查:先查实例缓存,避免每次都获取锁
  112. if name in self._instances:
  113. return self._instances[name]
  114. with self._locks[name]:
  115. # 再次检查,防止两个线程同时通过第一次检查
  116. if name not in self._instances:
  117. self._instances[name] = self._factories[name]()
  118. return self._instances[name]
  119. # ------------------------------------------------------------- introspection
  120. def is_registered(self, name):
  121. """name 是否已注册 factory。"""
  122. return name in self._factories
  123. def registered_names(self):
  124. """返回已注册的 predictor name 列表(副本,不暴露内部 dict)。"""
  125. with self._global_lock:
  126. return list(self._factories.keys())
  127. def is_loaded(self, name):
  128. """name 对应的实例是否已创建(用于启动预热检查)。"""
  129. return name in self._instances
  130. # ------------------------------------------------------------- lifecycle
  131. def reset(self, name):
  132. """清除单个 predictor 的实例缓存,保留 factory。
  133. 下次 ``get(name)`` 会重新调用 factory 创建新实例。
  134. 用于测试或模型热更新。
  135. """
  136. with self._global_lock:
  137. if name in self._instances:
  138. del self._instances[name]
  139. def clear(self):
  140. """清空所有 factory 和实例。主要用于测试。"""
  141. with self._global_lock:
  142. self._factories.clear()
  143. self._instances.clear()
  144. self._locks.clear()
  145. def warmup(self, names=None):
  146. """预热:提前创建实例。
  147. :param names: 指定预热的 name 列表,None 表示全部已注册的。
  148. 部署时可调用此方法避免首次请求慢。
  149. """
  150. targets = names if names is not None else self.registered_names()
  151. for n in targets:
  152. if self.is_registered(n) and not self.is_loaded(n):
  153. self.get(n)
  154. # =====================================================================
  155. # 默认全局注册表
  156. # =====================================================================
  157. _default_registry = PredictorRegistry()
  158. _defaults_registered = False
  159. _defaults_lock = RLock()
  160. def get_default_registry():
  161. """返回进程级默认 PredictorRegistry 实例。"""
  162. return _default_registry
  163. def is_defaults_registered():
  164. """``register_defaults`` 是否已执行过。用于测试断言。"""
  165. return _defaults_registered
  166. def register_defaults():
  167. """把原 ``interface/predictor.py`` 中 28 个 Predictor 注册到默认 registry。
  168. 工厂函数使用延迟 import,避免在 registry 模块加载时触发
  169. ``interface/predictor.py`` 加载(会拉起 TF/PyTorch 等重型依赖)。
  170. 幂等:多次调用只会注册一次。Phase 5 迁移单个 Predictor 时,
  171. 可在迁移后的模块中直接调用 ``registry.register(name, factory)`` 覆盖。
  172. Predictor 类与 key 的对应关系(见 DEFAULT_PREDICTOR_NAMES):
  173. codeName -> CodeNamePredict(config=sess_config)
  174. prem -> PREMPredict(config=sess_config)
  175. epc -> EPCPredict(config=sess_config)
  176. roleRule -> RoleRulePredictor()
  177. roleRuleFinal -> RoleRuleFinalAdd()
  178. tendereeRuleRecall -> TendereeRuleRecall()
  179. form -> FormPredictor(config=sess_config)
  180. time -> TimePredictor(config=sess_config)
  181. punish -> Punish_Extract()
  182. product -> ProductPredictor(config=sess_config)
  183. product_attrs -> ProductAttributesPredictor()
  184. channel -> DocChannel(config=sess_config)
  185. deposit_payment_way-> DepositPaymentWay()
  186. total_unit_money -> TotalUnitMoney()
  187. industry -> IndustryPredictor()
  188. rolegrade -> RoleGrade()
  189. moneygrade -> MoneyGrade()
  190. district -> DistrictPredictor()
  191. tableprem -> TablePremExtractor()
  192. candidate -> CandidateExtractor()
  193. websource_tenderee -> WebsourceTenderee()
  194. project_label -> ProjectLabel()
  195. industry_label -> IndustryLabel()
  196. pb_extract -> PBPredictor()
  197. property_label -> PropertyLabel()
  198. approval -> ApprovalPredictor()
  199. bid_score -> BiddingScore()
  200. entity_type_rule -> EntityTypeRulePredictor()
  201. """
  202. global _defaults_registered
  203. if _defaults_registered:
  204. return
  205. with _defaults_lock:
  206. if _defaults_registered:
  207. return
  208. # 延迟 import,避免循环依赖:
  209. # - register_defaults() 在 interface/predictor.py::getPredictor() 内被调用
  210. # - 此时 interface/predictor.py 已完成模块加载
  211. # - 这里只引用模块对象,不立即取属性,factory 执行时才取
  212. # - 使用绝对 import 与本仓库约定一致(from BiddingKG.dl.xxx)
  213. from BiddingKG.dl.interface import predictor as _p
  214. def _make(key, attr, use_sess_config=False):
  215. """构造延迟绑定的 factory。
  216. :param key: registry key(用于错误信息)。
  217. :param attr: interface.predictor 模块上的属性名(类名)。
  218. :param use_sess_config: 是否传 config=sess_config 参数。
  219. 原实现中模型类传 sess_config,规则类不传。
  220. """
  221. def _factory():
  222. cls = getattr(_p, attr)
  223. if use_sess_config:
  224. return cls(config=_p.sess_config)
  225. return cls()
  226. _factory.__name__ = "make_%s" % key
  227. return _factory
  228. # 按原 getPredictor() 的构造方式分类:
  229. # 1) 传 config=sess_config 的(模型类)
  230. _sess_config_keys = {
  231. "codeName": "CodeNamePredict",
  232. "prem": "PREMPredict",
  233. "epc": "EPCPredict",
  234. "form": "FormPredictor",
  235. "time": "TimePredictor",
  236. "product": "ProductPredictor",
  237. "channel": "DocChannel",
  238. }
  239. # 2) 不传参数的(规则类 / 已自管理配置的类)
  240. _no_arg_keys = {
  241. "roleRule": "RoleRulePredictor",
  242. "roleRuleFinal": "RoleRuleFinalAdd",
  243. "tendereeRuleRecall": "TendereeRuleRecall",
  244. "punish": "Punish_Extract",
  245. "product_attrs": "ProductAttributesPredictor",
  246. "deposit_payment_way": "DepositPaymentWay",
  247. "total_unit_money": "TotalUnitMoney",
  248. "industry": "IndustryPredictor",
  249. "rolegrade": "RoleGrade",
  250. "moneygrade": "MoneyGrade",
  251. "district": "DistrictPredictor",
  252. "tableprem": "TablePremExtractor",
  253. "candidate": "CandidateExtractor",
  254. "websource_tenderee": "WebsourceTenderee",
  255. "project_label": "ProjectLabel",
  256. "industry_label": "IndustryLabel",
  257. "pb_extract": "PBPredictor",
  258. "property_label": "PropertyLabel",
  259. "approval": "ApprovalPredictor",
  260. "bid_score": "BiddingScore",
  261. "entity_type_rule": "EntityTypeRulePredictor",
  262. }
  263. for key, attr in _sess_config_keys.items():
  264. _default_registry.register(key, _make(key, attr, use_sess_config=True))
  265. for key, attr in _no_arg_keys.items():
  266. _default_registry.register(key, _make(key, attr, use_sess_config=False))
  267. _defaults_registered = True