# -*- 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/.py - 每个 .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