pipeline.py 6.6 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184
  1. # -*- coding: utf-8 -*-
  2. """Pipeline 主编排器。
  3. 按 ARCHITECTURE.md Phase 8 设计,统一同步模式与队列模式的调用入口。
  4. ``api/single_server.py`` 和 ``api/queue_worker.py`` 都调用 ``Pipeline.run(ctx)``,
  5. 确保两条路径走完全相同的业务逻辑。
  6. 编排步骤(原 extract.py / run_single_server.py / run_model_server.py 的共有流程)::
  7. 1. preprocess — Preprocessing.get_preprocessed
  8. 2. codename — CodeNamePredict
  9. 3. prem — PREMPredict
  10. 4. rule — RoleRulePredictor
  11. 5. epc — EPCPredict
  12. 6. entity_link — entityLink.link_entitys
  13. 7. assemble — getAttributes.getPREMs
  14. 8. union — Preprocessing.union_result
  15. 使用方式::
  16. from BiddingKG.dl.pipeline.pipeline import Pipeline
  17. from BiddingKG.dl.pipeline.context import PipelineContext
  18. ctx = PipelineContext(doc_id="123", html=content, title=title)
  19. pipeline = Pipeline()
  20. result = pipeline.run(ctx)
  21. # result == ctx.result
  22. 约束(ARCHITECTURE.md §4.1):
  23. pipeline 只负责编排,不写复杂抽取逻辑。
  24. """
  25. from __future__ import absolute_import
  26. import time
  27. import logging
  28. from BiddingKG.dl.pipeline.context import PipelineContext
  29. logger = logging.getLogger(__name__)
  30. class Pipeline(object):
  31. """Pipeline 主编排器。
  32. 同步模式(single_server)和队列模式(queue_worker)共用此类。
  33. 模型实例在 ``__init__`` 时惰性创建(首次调用 ``run()`` 时才加载),
  34. 避免导入时的副作用。
  35. """
  36. def __init__(self):
  37. self._predictors = None
  38. self._preprocessing = None
  39. self._getAttributes = None
  40. self._entityLink = None
  41. # ------------------------------------------------------------------------
  42. # 惰性加载业务模块
  43. # ------------------------------------------------------------------------
  44. def _ensure_loaded(self):
  45. """惰性加载模型和业务模块(仅首次调用时)。"""
  46. if self._predictors is not None:
  47. return
  48. logger.info("Pipeline: loading models and modules...")
  49. import BiddingKG.dl.interface.predictor as predictor
  50. import BiddingKG.dl.interface.Preprocessing as Preprocessing
  51. import BiddingKG.dl.interface.getAttributes as getAttributes
  52. import BiddingKG.dl.entityLink.entityLink as entityLink
  53. self._predictors = {
  54. "codeName": predictor.CodeNamePredict(),
  55. "prem": predictor.PREMPredict(),
  56. "epc": predictor.EPCPredict(),
  57. "roleRule": predictor.RoleRulePredictor(),
  58. }
  59. self._preprocessing = Preprocessing
  60. self._getAttributes = getAttributes
  61. self._entityLink = entityLink
  62. logger.info("Pipeline: models loaded")
  63. # ------------------------------------------------------------------------
  64. # 主入口
  65. # ------------------------------------------------------------------------
  66. def run(self, ctx):
  67. """执行完整 Pipeline。
  68. :param ctx: PipelineContext,输入字段(doc_id/html/title)必须已填充
  69. :return: dict — 最终结果(同时写入 ctx.result)
  70. """
  71. self._ensure_loaded()
  72. t_total = time.time()
  73. # ---- 1. 预处理 ----
  74. t0 = time.time()
  75. k = ctx.doc_id or "pipeline"
  76. content = ctx.html
  77. ContentIDs = [[k, content, time.time(), ctx.doc_id, ctx.title]]
  78. list_articles, list_sentences, list_entitys, _cost_time = \
  79. self._preprocessing.get_preprocessed(ContentIDs, useselffool=True)
  80. ctx.articles = list_articles
  81. ctx.sentences = list_sentences
  82. ctx.entities = list_entitys
  83. ctx.cost_time["preprocess"] = time.time() - t0
  84. ctx.cost_time.update(_cost_time)
  85. # ---- 2-7. 预测 + 装配 ----
  86. self._run_predict(ctx, list_articles, list_sentences, list_entitys)
  87. # ---- 8. 合并结果 ----
  88. t0 = time.time()
  89. codeName = ctx.model_results.get("codeName")
  90. prem = ctx.model_results.get("prem")
  91. data = self._preprocessing.union_result(codeName, prem)[0][1]
  92. data["cost_time"] = ctx.cost_time
  93. data["success"] = True
  94. ctx.result = data
  95. ctx.cost_time["union"] = time.time() - t0
  96. ctx.cost_time["total"] = time.time() - t_total
  97. logger.info("Pipeline done for doc_id=%s, total=%.2fs" % (
  98. ctx.doc_id, ctx.cost_time["total"]))
  99. return ctx.result
  100. def run_predict(self, ctx, list_articles, list_sentences, list_entitys):
  101. """仅执行预测+装配(跳过预处理)。
  102. 用于队列模式中预处理已由其他 worker 完成的场景。
  103. """
  104. self._ensure_loaded()
  105. ctx.articles = list_articles
  106. ctx.sentences = list_sentences
  107. ctx.entities = list_entitys
  108. self._run_predict(ctx, list_articles, list_sentences, list_entitys)
  109. # 合并结果
  110. codeName = ctx.model_results.get("codeName")
  111. prem = ctx.model_results.get("prem")
  112. data = self._preprocessing.union_result(codeName, prem)[0][1]
  113. data["cost_time"] = ctx.cost_time
  114. data["success"] = True
  115. ctx.result = data
  116. return ctx.result
  117. # ------------------------------------------------------------------------
  118. # 内部:预测步骤
  119. # ------------------------------------------------------------------------
  120. def _run_predict(self, ctx, list_articles, list_sentences, list_entitys):
  121. """执行步骤 2-7:codename → prem → rule → epc → entityLink → assemble。"""
  122. p = self._predictors
  123. # 2. 项目编号/名称
  124. t0 = time.time()
  125. codeName = p["codeName"].predict(list_articles)
  126. ctx.model_results["codeName"] = codeName
  127. ctx.cost_time["codename"] = time.time() - t0
  128. # 3. 角色/金额候选
  129. t0 = time.time()
  130. p["prem"].predict(list_sentences, list_entitys)
  131. ctx.cost_time["prem"] = time.time() - t0
  132. # 4. 规则修正
  133. t0 = time.time()
  134. p["roleRule"].predict(list_articles, list_sentences, list_entitys, codeName)
  135. ctx.cost_time["rule"] = time.time() - t0
  136. # 5. 人员/EPC
  137. t0 = time.time()
  138. p["epc"].predict(list_sentences, list_entitys)
  139. ctx.cost_time["person"] = time.time() - t0
  140. # 6. 实体链接
  141. t0 = time.time()
  142. self._entityLink.link_entitys(list_entitys)
  143. ctx.cost_time["entity_link"] = time.time() - t0
  144. # 7. 业务装配
  145. t0 = time.time()
  146. prem = self._getAttributes.getPREMs(list_sentences, list_entitys, list_articles)
  147. ctx.model_results["prem"] = prem
  148. ctx.cost_time["attrs"] = time.time() - t0