product.py 13 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223
  1. # -*- coding: utf-8 -*-
  2. """``ProductPredictor`` — 产品/失败原因模型。
  3. Phase 5 从 ``interface/predictor.py``(约 3133-3330 行)迁出。
  4. 原 ``from common.Utils import *`` / ``from common.nerUtils import *``
  5. 已替换为显式 import;``os.path.dirname(__file__)`` 路径引用替换为
  6. ``predictors._common.INTERFACE_DIR``。
  7. """
  8. from __future__ import absolute_import
  9. import os
  10. import re
  11. import json
  12. import requests
  13. import numpy as np
  14. import tensorflow as tf
  15. from keras.preprocessing.sequence import pad_sequences
  16. from BiddingKG.dl.common.Utils import load
  17. from BiddingKG.dl.model_runtime.viterbi import viterbi_decode
  18. from BiddingKG.dl.interface.Entitys import Entity
  19. from BiddingKG.dl.services.external_inference.paieas import USE_API, API_URL
  20. from BiddingKG.dl.predictors._common import INTERFACE_DIR
  21. __all__ = ["ProductPredictor"]
  22. class ProductPredictor():
  23. def __init__(self,config=None):
  24. vocabpath = os.path.join(INTERFACE_DIR, "codename_vocab.pk")
  25. self.vocab = load(vocabpath)
  26. self.word2index = dict((w, i) for i, w in enumerate(np.array(self.vocab)))
  27. self.sess = tf.Session(graph=tf.Graph(),config=config)
  28. self.load_model()
  29. def load_model(self):
  30. # model_path = os.path.dirname(__file__)+'/product_savedmodel/product.pb'
  31. model_path = os.path.join(INTERFACE_DIR, "product_savedmodel", "productAndfailreason.pb")
  32. with self.sess.as_default():
  33. with self.sess.graph.as_default():
  34. output_graph_def = tf.GraphDef()
  35. with open(model_path, 'rb') as f:
  36. output_graph_def.ParseFromString(f.read())
  37. tf.import_graph_def(output_graph_def, name='')
  38. self.sess.run(tf.global_variables_initializer())
  39. self.char_input = self.sess.graph.get_tensor_by_name('CharInputs:0')
  40. self.length = self.sess.graph.get_tensor_by_name("Sum:0")
  41. self.dropout = self.sess.graph.get_tensor_by_name("Dropout:0")
  42. self.logit = self.sess.graph.get_tensor_by_name("logits/Reshape:0")
  43. self.tran = self.sess.graph.get_tensor_by_name("crf_loss/transitions:0")
  44. def decode(self,logits, lengths, matrix):
  45. paths = []
  46. small = -1000.0
  47. # start = np.asarray([[small] * 4 + [0]])
  48. start = np.asarray([[small]*7+[0]])
  49. for score, length in zip(logits, lengths):
  50. score = score[:length]
  51. pad = small * np.ones([length, 1])
  52. logits = np.concatenate([score, pad], axis=1)
  53. logits = np.concatenate([start, logits], axis=0)
  54. path, _ = viterbi_decode(logits, matrix)
  55. paths.append(path[1:])
  56. return paths
  57. def predict(self, list_sentences,list_entitys=None,list_articles=[], fail=False, MAX_AREA=5000, out_lines=[]):
  58. '''
  59. 预测实体代码,每个句子最多取MAX_AREA个字,超过截断
  60. :param list_sentences: 多篇公告句子列表,[[一篇公告句子列表],[公告句子列表]]
  61. :param list_entitys: 多篇公告实体列表
  62. :param MAX_AREA: 每个句子最多截取多少字
  63. :return: 把预测出来的实体放进实体类
  64. '''
  65. p = "(采购需求|需求分析|项目说明|(采购|合同|招标|询比?价|项目|服务|工程|标的|需求|建设|分包)(的?(主要|简要|基本|具体|名称及))?" \
  66. "(内容|概况|概述|范围|信息|规模|简介|介绍|说明|摘要|情况|名称)([及与和]((其它|\w{,2})[要需]求|发包范围|数量))?" \
  67. "|招标项目技术要求|服务要求|服务需求|项目目标|需求内容如下|建设规模|(设备|材料|仪器|需求|产品|采购单?)(清单|名称|信息))为?([::,]|$)"
  68. # sentence_range = [] #20240827 取消,修复线上接口产品耗时长问题
  69. # if len(out_lines) >= 3: # 三个以上大纲
  70. # for i in range(len(out_lines)-1):
  71. # text, s1, b1 = out_lines[i]
  72. # _, s2, b2 = out_lines[i+1]
  73. # if 3<text.find(':')<20:
  74. # text = text.split(':')[0]
  75. # if re.search(p, text[:15]):
  76. # sentence_range.append((s1, s2))
  77. with self.sess.as_default() as sess:
  78. with self.sess.graph.as_default():
  79. result = []
  80. product_list = []
  81. if fail and list_articles!=[]:
  82. text_list = [list_articles[0].content[:MAX_AREA]]
  83. chars = [[self.word2index.get(it, self.word2index.get('<unk>')) for it in text] for text in text_list]
  84. if USE_API:
  85. requests_result = requests.post(API_URL + "/predict_product",
  86. json={"inputs": chars}, verify=True)
  87. batch_paths = json.loads(requests_result.text)['result']
  88. lengths = json.loads(requests_result.text)['lengths']
  89. else:
  90. lengths, scores, tran_ = sess.run([self.length, self.logit, self.tran],
  91. feed_dict={
  92. self.char_input: np.asarray(chars),
  93. self.dropout: 1.0
  94. })
  95. batch_paths = self.decode(scores, lengths, tran_)
  96. for text, path, length in zip(text_list, batch_paths, lengths):
  97. tags = ''.join([str(it) for it in path[:length]])
  98. # 提取产品
  99. for it in re.finditer("12*3", tags):
  100. start = it.start()
  101. end = it.end()
  102. if text[start:end] in ['服务', '运费']:
  103. continue
  104. _entity = Entity(doc_id=list_articles[0].id, entity_id="%s_%s_%s_%s" % (
  105. list_articles[0].doc_id, 0, start, end),
  106. entity_text=text[start:end],
  107. entity_type="product", sentence_index=0,
  108. begin_index=0, end_index=0, wordOffset_begin=start,
  109. wordOffset_end=end)
  110. list_entitys[0].append(_entity)
  111. product_list.append(text[start:end])
  112. # 提取失败原因
  113. for it in re.finditer("45*6", tags):
  114. start = it.start()
  115. end = it.end()
  116. result.append(text[start:end].replace('?', '').strip())
  117. reasons = []
  118. for it in result:
  119. if "(√)" in it or "(√)" in it:
  120. reasons = [it]
  121. break
  122. if reasons != [] and (it not in reasons[-1] and it not in reasons):
  123. reasons.append(it)
  124. elif reasons == []:
  125. reasons.append(it)
  126. if reasons == []: # 如果模型识别不到失败原因 就用规则补充
  127. for text in text_list:
  128. ser1 = re.search('\w{,4}(理由|原因):\s*((第\d+包|标项\d+|原因类型)?[::]?[\s*\w,]{2,30}((不满?足|少于|未达)((法定)?[123一二三两]家|(规定)?要求)|(项目|采购)(终止|废标)),?)+',text)
  129. ser2 = re.search(
  130. '\w{,4}(理由|原因):\s*(第\d+包|标项\d+|原因类型)?[::]?[\s*\w]{4,30},', text)
  131. if ser1:
  132. reasons.append(ser1.group(0))
  133. break
  134. elif ser2:
  135. reasons.append(ser2.group(0))
  136. break
  137. return {'fail_reason':';'.join(reasons)}, product_list
  138. if list_entitys is None:
  139. list_entitys = [[] for _ in range(len(list_sentences))]
  140. for list_sentence, list_entity in zip(list_sentences,list_entitys):
  141. if len(list_sentence)==0:
  142. result.append({"product":[]})
  143. continue
  144. # 20240827 取消,修复线上接口产品耗时长问题
  145. # if sentence_range: # 20240815 如果有招标内容大纲,只从前两句及大纲内提取产品,避免类似 514920213 提取错其他内容 银行流水
  146. # new_list = []
  147. # word_num = 0
  148. # for sentence in list_sentence:
  149. # if sentence.sentence_index<2:
  150. # new_list.append(sentence)
  151. # continue
  152. # for s1, s2 in sentence_range:
  153. # if sentence.sentence_index < s1:
  154. # continue
  155. # elif s1<=sentence.sentence_index <=s2:
  156. # new_list.append(sentence)
  157. # word_num += len(sentence.sentence_text)
  158. # elif sentence.sentence_index >= s2:
  159. # break
  160. # if word_num > 100:
  161. # list_sentence = new_list
  162. list_sentence.sort(key=lambda x:len(x.sentence_text), reverse=True)
  163. _begin_index = 0
  164. item = {"product":[]}
  165. temp_list = []
  166. while True:
  167. MAX_LEN = len(list_sentence[_begin_index].sentence_text)
  168. if MAX_LEN > MAX_AREA:
  169. MAX_LEN = MAX_AREA
  170. _LEN = MAX_AREA//MAX_LEN
  171. chars = [sentence.sentence_text[:MAX_LEN] for sentence in list_sentence[_begin_index:_begin_index+_LEN]]
  172. chars = [[self.word2index.get(it, self.word2index.get('<unk>')) for it in l] for l in chars]
  173. chars = pad_sequences(chars, maxlen=MAX_LEN, padding="post", truncating="post")
  174. if USE_API:
  175. requests_result = requests.post(API_URL + "/predict_product",
  176. json={"inputs": chars.tolist()}, verify=True)
  177. batch_paths = json.loads(requests_result.text)['result']
  178. lengths = json.loads(requests_result.text)['lengths']
  179. else:
  180. lengths, scores, tran_ = sess.run([self.length, self.logit, self.tran],
  181. feed_dict={
  182. self.char_input: np.asarray(chars),
  183. self.dropout: 1.0
  184. })
  185. batch_paths = self.decode(scores, lengths, tran_)
  186. for sentence, path, length in zip(list_sentence[_begin_index:_begin_index+_LEN],batch_paths, lengths):
  187. tags = ''.join([str(it) for it in path[:length]])
  188. for it in re.finditer("12*3", tags):
  189. start = it.start()
  190. end = it.end()
  191. if sentence.sentence_text[start:end] in ['服务', '运费']: # 20260522修复 739917393 去除此类型产品
  192. continue
  193. _entity = Entity(doc_id=sentence.doc_id, entity_id="%s_%s_%s_%s" % (
  194. sentence.doc_id, sentence.sentence_index, start, end),
  195. entity_text=sentence.sentence_text[start:end],
  196. entity_type="product", sentence_index=sentence.sentence_index,
  197. begin_index=0, end_index=0, wordOffset_begin=start,
  198. wordOffset_end=end,in_attachment=sentence.in_attachment)
  199. list_entity.append(_entity)
  200. temp_list.append(sentence.sentence_text[start:end])
  201. product_list.append(sentence.sentence_text[start:end])
  202. # item["product"] = list(set(temp_list))
  203. # result.append(item)
  204. if _begin_index+_LEN >= len(list_sentence):
  205. break
  206. _begin_index += _LEN
  207. item["product"] = list(set(temp_list))
  208. result.append(item) # 修正bug
  209. return {'fail_reason': ""},product_list