| 123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302303304305306307308309310311312313314315316317318319320321322323324325326327328329330331332333334335336337338339340341342343344345346347348349350351352353354355356357358359360361362363364365366367368369370371372373374375376377378379380381382383384385386 |
- # -*- coding: utf-8 -*-
- """RoleRuleEngine — YAML 驱动的角色/金额模型后修正引擎(角色分类流程优化 Phase C Stage 2)。
- 职责
- ====
- 消费 ``dl/rules/patterns/role_context_fix.yaml``(ROLE)与
- ``role_money_fix.yaml``(MONEY)中带 ``stage`` 字段的机器可读修正规则,
- 按「阶段序 + 阶段内首条命中」的语义执行,精确复刻原
- ``prem.predict_role`` / ``predict_money`` / ``correct_money_by_rule``
- 中内联 if/elif 修正链的行为(行为等价重构,Phase D 双跑回归验证)。
- 阶段序(与原代码 if/elif 短路语义一一对应)
- ============================================
- 角色(correct_role)::
- global_pre(非终止:命中后 label 改写继续走后续判定)
- → 阈值过滤(label∈[0,4] 且 values[label] < P_MODEL_THRESHOLD → label=5,终止)
- → seq_check(label∈[2,3,4] 命中 → label=5,终止)
- → 分支:label=0→l0;label=2→l2;label=1→l1;label∈[3,4]→l34;
- label=5→l5_win_yes → notify → l5(逐阶段首条命中即止)
- 金额(correct_money)::
- 阈值过滤(label∈[0,1] 且 values[label] < P_MODEL_THRESHOLD → label=2,终止)
- → 分支:label=1→m1;label=0→m0;label=2→m_bid
- 金额标题类别(correct_money_by_doc)::
- doc_title 阶段独立执行(entity_filter 过滤实体,不依赖上下文切片),
- 替代原 correct_money_by_rule。
- 依赖方向说明
- ============
- 本模块位于 CORE(predictors),只读消费 ``dl/rules`` 的 patterns 运行时数据
- 与 loader 框架,不依赖 ``dl/rules/generated/``(ARCHITECTURE.md §4.3)。
- ``is_agency`` 谓词来自同包 ``predictors/_common.py``(原 prem.py 即如此使用)。
- 规则条件 DSL 见 ``dl/rules/schema/pattern_schema.json`` definitions.condition。
- """
- from __future__ import absolute_import
- import re
- from typing import Any, Dict, List, Optional, Tuple
- from BiddingKG.dl.predictors._common import is_agency
- from BiddingKG.dl.predictors.role_context import P_MODEL_THRESHOLD
- from BiddingKG.dl.rules.loader import RuleLoader
- __all__ = ["RoleRuleEngine"]
- #: 角色分支阶段调度:当前 label → 依次尝试的阶段列表
- _ROLE_BRANCH_STAGES = {
- 0: ("l0",),
- 2: ("l2",),
- 1: ("l1",),
- 3: ("l34",),
- 4: ("l34",),
- 5: ("l5_win_yes", "notify", "l5"),
- }
- #: 金额分支阶段调度:模型原始 label(阈值过滤后)→ 分支阶段
- _MONEY_BRANCH_STAGES = {
- 1: ("m1",),
- 0: ("m0",),
- 2: ("m_bid",),
- }
- #: 条件求值缺省的 title_content_head / content_head 截断长度
- _DEFAULT_HEAD_SIZE = 100
- class RoleRuleEngine(object):
- """YAML 驱动的角色/金额模型后修正引擎。
- 用法(prem.py)::
- engine = RoleRuleEngine()
- label = engine.correct_role(entity, label, values, front, middle, behind)
- entity.set_Role(label, values)
- values 被就地修改(与原代码一致);label 由返回值带回。
- """
- def __init__(self, agency_checker=None):
- """
- :param agency_checker: is_agency 谓词(可注入替换,缺省用
- ``predictors/_common.is_agency``,与原 prem.py 行为一致)
- """
- self._agency_checker = agency_checker or is_agency
- self._role_stages = self._group_stages("ROLE")
- self._money_stages = self._group_stages("MONEY")
- # ------------------------------------------------------------------
- # 规则装载
- # ------------------------------------------------------------------
- @staticmethod
- def _group_stages(category):
- """按 stage 分组规则(组内保持 YAML 顺序)。"""
- stages = {}
- for rule in RuleLoader.get_rules(category=category, staged=True):
- stages.setdefault(rule["stage"], []).append(rule)
- return stages
- # ------------------------------------------------------------------
- # 角色:替代 prem.predict_role 内联修正链
- # ------------------------------------------------------------------
- def correct_role(self, entity, label, values, front, middle, behind):
- """执行角色模型后修正(行为等价原 prem.predict_role if/elif 链)。
- :param entity: 实体(entity_text 用于 is_agency 谓词)
- :param label: 模型预测 label(int)
- :param values: 模型输出概率数组(就地修改)
- :param front/middle/behind: role_model_text 三元组(前文23/实体/后文25)
- :return: 修正后的 label
- """
- env = {
- "front": front,
- "behind": behind,
- "whole": front[-10:] + middle + behind[:10],
- "entity_text": entity.entity_text,
- }
- # ---- global_pre:非终止,命中后以新 label 继续走阈值/分支 ----
- rule = self._first_match("global_pre", label, env, entity, values)
- if rule is not None:
- label = self._apply(rule, label, values)
- # ---- 阈值过滤(程序逻辑,终止)----
- if label in (0, 1, 2, 3, 4) and values[label] < P_MODEL_THRESHOLD:
- return 5
- # ---- seq_check(终止)----
- if label in (2, 3, 4):
- rule = self._first_match("seq_check", label, env, entity, values)
- if rule is not None:
- return 5
- # ---- 分支阶段 ----
- for stage in _ROLE_BRANCH_STAGES.get(label, ()):
- rule = self._first_match(stage, label, env, entity, values)
- if rule is not None:
- label = self._apply(rule, label, values)
- break
- return label
- # ------------------------------------------------------------------
- # 金额:替代 prem.predict_money 内联修正链
- # ------------------------------------------------------------------
- def correct_money(self, entity, label, values, front, middle, behind):
- """执行金额模型后修正(行为等价原 prem.predict_money if/elif 链)。
- :param front/middle/behind: money_model_text 三元组(前文13/实体/后文15)
- :return: 修正后的 label
- """
- env = {
- "front": front,
- "behind": behind,
- "entity_text": entity.entity_text,
- }
- # ---- 阈值过滤(程序逻辑,终止)----
- if label in (0, 1) and values[label] < P_MODEL_THRESHOLD:
- return 2
- for stage in _MONEY_BRANCH_STAGES.get(label, ()):
- rule = self._first_match(stage, label, env, entity, values)
- if rule is not None:
- label = self._apply(rule, label, values)
- break
- return label
- # ------------------------------------------------------------------
- # 金额标题类别:替代 prem.correct_money_by_rule
- # ------------------------------------------------------------------
- def correct_money_by_doc(self, title, content, list_entitys):
- """按公告标题/正文类别批量修正金额实体(doc_title 阶段)。
- :param title: 公告标题
- :param content: 首篇文章正文(原代码 list_articles[0].content)
- :param list_entitys: 与原 correct_money_by_rule 相同的嵌套实体列表
- """
- rules = self._money_stages.get("doc_title", ())
- if not rules:
- return
- for list_entity in list_entitys:
- for entity in list_entity:
- env = {
- "title": title,
- "content": content,
- "entity_text": entity.entity_text,
- }
- for rule in rules:
- if not self._entity_filter_ok(entity, rule):
- continue
- if self._fires(rule, env, entity, entity.values):
- new_label = self._apply(rule, entity.label, entity.values)
- entity.set_Money(new_label, entity.values)
- break
- # ------------------------------------------------------------------
- # 内部:规则匹配与动作
- # ------------------------------------------------------------------
- def _first_match(self, stage, label, env, entity, values):
- """阶段内按 YAML 顺序找首条命中规则(对应原 elif 链短路)。"""
- for rule in list(self._role_stages.get(stage, ())) + list(self._money_stages.get(stage, ())):
- if "model_labels" in rule and label not in rule["model_labels"]:
- continue
- if self._fires(rule, env, entity, values):
- return rule
- return None
- def _fires(self, rule, env, entity, values):
- """规则触发判定:trigger(pattern 或 or_trigger)且 guard 且无 guard_not。"""
- # ---- trigger ----
- trig = False
- pattern = rule.get("pattern_compiled")
- if pattern is not None:
- text = self._apply_slice(
- rule.get("applies_to", "front"),
- rule.get("window", {}).get("size"),
- env,
- )
- if text is not None:
- method = rule.get("method", "search")
- if method == "match":
- trig = pattern.match(text) is not None
- elif method == "fullmatch":
- trig = pattern.fullmatch(text) is not None
- else:
- trig = pattern.search(text) is not None
- if not trig:
- for cond in rule.get("or_trigger", ()):
- if self._eval_cond(cond, env, entity, values):
- trig = True
- break
- if not trig:
- return False
- # ---- guard:AND ----
- for cond in rule.get("guard", ()):
- if not self._eval_cond(cond, env, entity, values):
- return False
- # ---- guard_not:OR 排除 ----
- for cond in rule.get("guard_not", ()):
- if self._eval_cond(cond, env, entity, values):
- return False
- return True
- @staticmethod
- def _apply(rule, label, values):
- """应用 set_label / set_prob,返回新 label。values 就地修改。"""
- new_label = rule.get("set_label", label)
- set_prob = rule.get("set_prob")
- if set_prob is not None:
- index = set_prob.get("index", rule.get("set_label", label))
- values[index] = set_prob["value"]
- return new_label
- # ------------------------------------------------------------------
- # 内部:切片与条件求值
- # ------------------------------------------------------------------
- @staticmethod
- def _apply_slice(on, size, env):
- """解析条件目标切片;front/behind 按 size 截断(front 取末段)。"""
- if on == "front":
- text = env.get("front")
- if text is None:
- return None
- return text[-size:] if size else text
- if on == "behind":
- text = env.get("behind")
- if text is None:
- return None
- return text[:size] if size else text
- if on == "whole":
- return env.get("whole")
- if on == "entity_text":
- return env.get("entity_text")
- if on == "title":
- return env.get("title")
- if on == "content":
- return env.get("content")
- if on == "content_head":
- content = env.get("content")
- if content is None:
- return None
- return content[: size or _DEFAULT_HEAD_SIZE]
- if on == "title_content_head":
- title = env.get("title")
- content = env.get("content")
- if title is None or content is None:
- return None
- return title + content[: size or _DEFAULT_HEAD_SIZE]
- return None
- def _eval_cond(self, cond, env, entity, values):
- """求值结构化条件(schema definitions.condition)。"""
- ctype = cond["type"]
- if ctype == "regex":
- text = self._apply_slice(cond.get("on", "front"), cond.get("size"), env)
- if text is None:
- return False
- return re_search(cond["pattern"], text) is not None
- if ctype == "regex_not":
- text = self._apply_slice(cond.get("on", "front"), cond.get("size"), env)
- if text is None:
- return False
- return re_search(cond["pattern"], text) is None
- if ctype == "match":
- text = self._apply_slice(cond.get("on", "front"), cond.get("size"), env)
- if text is None:
- return False
- return re_match(cond["pattern"], text) is not None
- if ctype == "is_agency":
- return bool(self._agency_checker(env.get("entity_text") or ""))
- if ctype == "not_is_agency":
- return not bool(self._agency_checker(env.get("entity_text") or ""))
- if ctype == "value_lt":
- return values[cond["index"]] < cond["threshold"]
- if ctype == "notes_in":
- return getattr(entity, "notes", None) in cond["values"]
- if ctype == "any":
- return any(self._eval_cond(c, env, entity, values) for c in cond["of"])
- if ctype == "all":
- return all(self._eval_cond(c, env, entity, values) for c in cond["of"])
- if ctype == "count_eq":
- text = self._apply_slice(cond.get("on", "title"), cond.get("size"), env)
- if text is None:
- return False
- return len(re_findall(cond["pattern"], text)) == cond["equals"]
- raise ValueError("未知条件类型: %s" % ctype)
- @staticmethod
- def _entity_filter_ok(entity, rule):
- """doc_title 阶段实体过滤(读规则的 entity_filter 字段)。"""
- flt = rule.get("entity_filter")
- if not flt:
- return True
- if "entity_type" in flt and entity.entity_type != flt["entity_type"]:
- return False
- if "notes_in" in flt and getattr(entity, "notes", None) not in flt["notes_in"]:
- return False
- if "label" in flt and entity.label != flt["label"]:
- return False
- return True
- # ----------------------------------------------------------------------------
- # 条件正则缓存(guard/match/count_eq 用,避免每次求值重复编译)
- # ----------------------------------------------------------------------------
- _COND_CACHE: Dict[str, Any] = {}
- def re_search(pattern, text):
- key = "s:" + pattern
- pat = _COND_CACHE.get(key)
- if pat is None:
- pat = _COND_CACHE[key] = re.compile(pattern)
- return pat.search(text)
- def re_match(pattern, text):
- key = "m:" + pattern
- pat = _COND_CACHE.get(key)
- if pat is None:
- pat = _COND_CACHE[key] = re.compile(pattern)
- return pat.match(text)
- def re_findall(pattern, text):
- key = "f:" + pattern
- pat = _COND_CACHE.get(key)
- if pat is None:
- pat = _COND_CACHE[key] = re.compile(pattern)
- return pat.findall(text)
|