| 123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302303304305306307308309310311312313314315316 |
- # -*- coding: utf-8 -*-
- """Predictor 注册表(PredictorRegistry)。
- 替代 interface/predictor.py 中的全局 ``dict_predictor`` 字典。
- 设计目标(ARCHITECTURE.md §4.3、§6.2、Phase 2):
- 1. 线程安全:每个 predictor 有独立 RLock,与原 dict_predictor 行为一致。
- 2. 懒加载:首次 ``get(name)`` 才调用 factory 创建实例,避免启动时全量加载。
- 3. 显式注册:通过 ``register(name, factory)`` 注册,便于 Phase 5 逐个迁移 Predictor。
- 4. 兼容旧入口:``interface/predictor.py::getPredictor()`` 改为委托本注册表,
- 老 import 路径无需改动。
- 5. 不强制 BasePredictor 子类:只要是可调用 factory 返回实例即可,
- 保证现有 28 个 Predictor 类无需立即改造。
- 迁移路径:
- Phase 2(本文件):
- - 所有 28 个 Predictor 仍在 interface/predictor.py 中定义
- - register_defaults() 用 lambda 延迟 import 旧位置类
- - getPredictor() 委托 get_default_registry().get(name)
- Phase 5:
- - Predictor 类逐个迁移到 predictors/<name>.py
- - 每个 <name>.py 在 import 时调用 register() 覆盖默认 factory
- - interface/predictor.py 变成纯 re-export 兼容层,最终删除
- 用法示例::
- from BiddingKG.dl.predictors.registry import get_default_registry, register_defaults
- register_defaults() # 幂等,多次调用只注册一次
- predictor = get_default_registry().get("codeName")
- # Phase 5 自定义注册:
- from BiddingKG.dl.predictors.codename import CodeNamePredict
- registry = get_default_registry()
- registry.register("codeName", lambda: CodeNamePredict()) # 覆盖默认 factory
- """
- from __future__ import absolute_import
- from threading import RLock
- __all__ = [
- "PredictorRegistry",
- "get_default_registry",
- "register_defaults",
- "is_defaults_registered",
- ]
- #: 原 dict_predictor 中 28 个 predictor 的 key 列表(按原文件顺序)。
- #: 顺序仅用于文档展示,registry 内部用 dict 存储不依赖顺序。
- DEFAULT_PREDICTOR_NAMES = (
- "codeName",
- "prem",
- "epc",
- "roleRule",
- "roleRuleFinal",
- "tendereeRuleRecall",
- "form",
- "time",
- "punish",
- "product",
- "product_attrs",
- "channel",
- "deposit_payment_way",
- "total_unit_money",
- "industry",
- "rolegrade",
- "moneygrade",
- "district",
- "tableprem",
- "candidate",
- "websource_tenderee",
- "project_label",
- "industry_label",
- "pb_extract",
- "property_label",
- "approval",
- "bid_score",
- "entity_type_rule",
- )
- class PredictorRegistry(object):
- """线程安全的 Predictor 注册表。
- 与原 ``dict_predictor`` 的对应关系:
- - ``dict_predictor[name]["Lock"]`` -> 内部 ``self._locks[name]``
- - ``dict_predictor[name]["predictor"]`` -> 内部 ``self._instances[name]``
- - ``getPredictor(name)`` -> ``self.get(name)``
- factory 是零参数 callable,返回 Predictor 实例。
- factory 在 ``register()`` 时不被调用,只在首次 ``get()`` 时调用一次。
- """
- def __init__(self):
- self._factories = {}
- self._instances = {}
- self._locks = {}
- # 保护 _factories / _locks 字典本身的结构性变更(register/clear/reset)。
- # 单个 predictor 实例的创建由 self._locks[name] 保护。
- self._global_lock = RLock()
- # ------------------------------------------------------------------ register
- def register(self, name, factory):
- """注册或覆盖一个 predictor factory。
- :param name: predictor key,例如 "codeName"。
- :param factory: 零参数 callable,返回 predictor 实例。
- 重复 register 同名 key 会覆盖旧 factory,但已创建的实例不会立即失效,
- 需调用 :meth:`reset` 清除旧实例。
- 幂等性:重复 register 同一个 (name, factory) 不影响已创建实例。
- """
- if not callable(factory):
- raise TypeError("factory must be callable, got %r" % type(factory))
- with self._global_lock:
- self._factories[name] = factory
- if name not in self._locks:
- self._locks[name] = RLock()
- # --------------------------------------------------------------------- get
- def get(self, name):
- """获取 predictor 实例,首次调用时触发 factory 创建。
- :raises KeyError: name 未注册时抛出。
- 原实现抛 ``NameError("no this type of predictor")``,
- ``interface/predictor.py::getPredictor`` 会捕获并转换为 NameError
- 以保持兼容。
- """
- if name not in self._factories:
- raise KeyError("no predictor registered for: %s" % name)
- # 双重检查:先查实例缓存,避免每次都获取锁
- if name in self._instances:
- return self._instances[name]
- with self._locks[name]:
- # 再次检查,防止两个线程同时通过第一次检查
- if name not in self._instances:
- self._instances[name] = self._factories[name]()
- return self._instances[name]
- # ------------------------------------------------------------- introspection
- def is_registered(self, name):
- """name 是否已注册 factory。"""
- return name in self._factories
- def registered_names(self):
- """返回已注册的 predictor name 列表(副本,不暴露内部 dict)。"""
- with self._global_lock:
- return list(self._factories.keys())
- def is_loaded(self, name):
- """name 对应的实例是否已创建(用于启动预热检查)。"""
- return name in self._instances
- # ------------------------------------------------------------- lifecycle
- def reset(self, name):
- """清除单个 predictor 的实例缓存,保留 factory。
- 下次 ``get(name)`` 会重新调用 factory 创建新实例。
- 用于测试或模型热更新。
- """
- with self._global_lock:
- if name in self._instances:
- del self._instances[name]
- def clear(self):
- """清空所有 factory 和实例。主要用于测试。"""
- with self._global_lock:
- self._factories.clear()
- self._instances.clear()
- self._locks.clear()
- def warmup(self, names=None):
- """预热:提前创建实例。
- :param names: 指定预热的 name 列表,None 表示全部已注册的。
- 部署时可调用此方法避免首次请求慢。
- """
- targets = names if names is not None else self.registered_names()
- for n in targets:
- if self.is_registered(n) and not self.is_loaded(n):
- self.get(n)
- # =====================================================================
- # 默认全局注册表
- # =====================================================================
- _default_registry = PredictorRegistry()
- _defaults_registered = False
- _defaults_lock = RLock()
- def get_default_registry():
- """返回进程级默认 PredictorRegistry 实例。"""
- return _default_registry
- def is_defaults_registered():
- """``register_defaults`` 是否已执行过。用于测试断言。"""
- return _defaults_registered
- def register_defaults():
- """把原 ``interface/predictor.py`` 中 28 个 Predictor 注册到默认 registry。
- 工厂函数使用延迟 import,避免在 registry 模块加载时触发
- ``interface/predictor.py`` 加载(会拉起 TF/PyTorch 等重型依赖)。
- 幂等:多次调用只会注册一次。Phase 5 迁移单个 Predictor 时,
- 可在迁移后的模块中直接调用 ``registry.register(name, factory)`` 覆盖。
- Predictor 类与 key 的对应关系(见 DEFAULT_PREDICTOR_NAMES):
- codeName -> CodeNamePredict(config=sess_config)
- prem -> PREMPredict(config=sess_config)
- epc -> EPCPredict(config=sess_config)
- roleRule -> RoleRulePredictor()
- roleRuleFinal -> RoleRuleFinalAdd()
- tendereeRuleRecall -> TendereeRuleRecall()
- form -> FormPredictor(config=sess_config)
- time -> TimePredictor(config=sess_config)
- punish -> Punish_Extract()
- product -> ProductPredictor(config=sess_config)
- product_attrs -> ProductAttributesPredictor()
- channel -> DocChannel(config=sess_config)
- deposit_payment_way-> DepositPaymentWay()
- total_unit_money -> TotalUnitMoney()
- industry -> IndustryPredictor()
- rolegrade -> RoleGrade()
- moneygrade -> MoneyGrade()
- district -> DistrictPredictor()
- tableprem -> TablePremExtractor()
- candidate -> CandidateExtractor()
- websource_tenderee -> WebsourceTenderee()
- project_label -> ProjectLabel()
- industry_label -> IndustryLabel()
- pb_extract -> PBPredictor()
- property_label -> PropertyLabel()
- approval -> ApprovalPredictor()
- bid_score -> BiddingScore()
- entity_type_rule -> EntityTypeRulePredictor()
- """
- global _defaults_registered
- if _defaults_registered:
- return
- with _defaults_lock:
- if _defaults_registered:
- return
- # 延迟 import,避免循环依赖:
- # - register_defaults() 在 interface/predictor.py::getPredictor() 内被调用
- # - 此时 interface/predictor.py 已完成模块加载
- # - 这里只引用模块对象,不立即取属性,factory 执行时才取
- # - 使用绝对 import 与本仓库约定一致(from BiddingKG.dl.xxx)
- from BiddingKG.dl.interface import predictor as _p
- def _make(key, attr, use_sess_config=False):
- """构造延迟绑定的 factory。
- :param key: registry key(用于错误信息)。
- :param attr: interface.predictor 模块上的属性名(类名)。
- :param use_sess_config: 是否传 config=sess_config 参数。
- 原实现中模型类传 sess_config,规则类不传。
- """
- def _factory():
- cls = getattr(_p, attr)
- if use_sess_config:
- return cls(config=_p.sess_config)
- return cls()
- _factory.__name__ = "make_%s" % key
- return _factory
- # 按原 getPredictor() 的构造方式分类:
- # 1) 传 config=sess_config 的(模型类)
- _sess_config_keys = {
- "codeName": "CodeNamePredict",
- "prem": "PREMPredict",
- "epc": "EPCPredict",
- "form": "FormPredictor",
- "time": "TimePredictor",
- "product": "ProductPredictor",
- "channel": "DocChannel",
- }
- # 2) 不传参数的(规则类 / 已自管理配置的类)
- _no_arg_keys = {
- "roleRule": "RoleRulePredictor",
- "roleRuleFinal": "RoleRuleFinalAdd",
- "tendereeRuleRecall": "TendereeRuleRecall",
- "punish": "Punish_Extract",
- "product_attrs": "ProductAttributesPredictor",
- "deposit_payment_way": "DepositPaymentWay",
- "total_unit_money": "TotalUnitMoney",
- "industry": "IndustryPredictor",
- "rolegrade": "RoleGrade",
- "moneygrade": "MoneyGrade",
- "district": "DistrictPredictor",
- "tableprem": "TablePremExtractor",
- "candidate": "CandidateExtractor",
- "websource_tenderee": "WebsourceTenderee",
- "project_label": "ProjectLabel",
- "industry_label": "IndustryLabel",
- "pb_extract": "PBPredictor",
- "property_label": "PropertyLabel",
- "approval": "ApprovalPredictor",
- "bid_score": "BiddingScore",
- "entity_type_rule": "EntityTypeRulePredictor",
- }
- for key, attr in _sess_config_keys.items():
- _default_registry.register(key, _make(key, attr, use_sess_config=True))
- for key, attr in _no_arg_keys.items():
- _default_registry.register(key, _make(key, attr, use_sess_config=False))
- _defaults_registered = True
|