| 123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223 |
- # -*- coding: utf-8 -*-
- """``ProductPredictor`` — 产品/失败原因模型。
- Phase 5 从 ``interface/predictor.py``(约 3133-3330 行)迁出。
- 原 ``from common.Utils import *`` / ``from common.nerUtils import *``
- 已替换为显式 import;``os.path.dirname(__file__)`` 路径引用替换为
- ``predictors._common.INTERFACE_DIR``。
- """
- from __future__ import absolute_import
- import os
- import re
- import json
- import requests
- import numpy as np
- import tensorflow as tf
- from keras.preprocessing.sequence import pad_sequences
- from BiddingKG.dl.common.Utils import load
- from BiddingKG.dl.model_runtime.viterbi import viterbi_decode
- from BiddingKG.dl.interface.Entitys import Entity
- from BiddingKG.dl.services.external_inference.paieas import USE_API, API_URL
- from BiddingKG.dl.predictors._common import INTERFACE_DIR
- __all__ = ["ProductPredictor"]
- class ProductPredictor():
- def __init__(self,config=None):
- vocabpath = os.path.join(INTERFACE_DIR, "codename_vocab.pk")
- self.vocab = load(vocabpath)
- self.word2index = dict((w, i) for i, w in enumerate(np.array(self.vocab)))
- self.sess = tf.Session(graph=tf.Graph(),config=config)
- self.load_model()
- def load_model(self):
- # model_path = os.path.dirname(__file__)+'/product_savedmodel/product.pb'
- model_path = os.path.join(INTERFACE_DIR, "product_savedmodel", "productAndfailreason.pb")
- with self.sess.as_default():
- with self.sess.graph.as_default():
- output_graph_def = tf.GraphDef()
- with open(model_path, 'rb') as f:
- output_graph_def.ParseFromString(f.read())
- tf.import_graph_def(output_graph_def, name='')
- self.sess.run(tf.global_variables_initializer())
- self.char_input = self.sess.graph.get_tensor_by_name('CharInputs:0')
- self.length = self.sess.graph.get_tensor_by_name("Sum:0")
- self.dropout = self.sess.graph.get_tensor_by_name("Dropout:0")
- self.logit = self.sess.graph.get_tensor_by_name("logits/Reshape:0")
- self.tran = self.sess.graph.get_tensor_by_name("crf_loss/transitions:0")
- def decode(self,logits, lengths, matrix):
- paths = []
- small = -1000.0
- # start = np.asarray([[small] * 4 + [0]])
- start = np.asarray([[small]*7+[0]])
- for score, length in zip(logits, lengths):
- score = score[:length]
- pad = small * np.ones([length, 1])
- logits = np.concatenate([score, pad], axis=1)
- logits = np.concatenate([start, logits], axis=0)
- path, _ = viterbi_decode(logits, matrix)
- paths.append(path[1:])
- return paths
- def predict(self, list_sentences,list_entitys=None,list_articles=[], fail=False, MAX_AREA=5000, out_lines=[]):
- '''
- 预测实体代码,每个句子最多取MAX_AREA个字,超过截断
- :param list_sentences: 多篇公告句子列表,[[一篇公告句子列表],[公告句子列表]]
- :param list_entitys: 多篇公告实体列表
- :param MAX_AREA: 每个句子最多截取多少字
- :return: 把预测出来的实体放进实体类
- '''
- p = "(采购需求|需求分析|项目说明|(采购|合同|招标|询比?价|项目|服务|工程|标的|需求|建设|分包)(的?(主要|简要|基本|具体|名称及))?" \
- "(内容|概况|概述|范围|信息|规模|简介|介绍|说明|摘要|情况|名称)([及与和]((其它|\w{,2})[要需]求|发包范围|数量))?" \
- "|招标项目技术要求|服务要求|服务需求|项目目标|需求内容如下|建设规模|(设备|材料|仪器|需求|产品|采购单?)(清单|名称|信息))为?([::,]|$)"
- # sentence_range = [] #20240827 取消,修复线上接口产品耗时长问题
- # if len(out_lines) >= 3: # 三个以上大纲
- # for i in range(len(out_lines)-1):
- # text, s1, b1 = out_lines[i]
- # _, s2, b2 = out_lines[i+1]
- # if 3<text.find(':')<20:
- # text = text.split(':')[0]
- # if re.search(p, text[:15]):
- # sentence_range.append((s1, s2))
- with self.sess.as_default() as sess:
- with self.sess.graph.as_default():
- result = []
- product_list = []
- if fail and list_articles!=[]:
- text_list = [list_articles[0].content[:MAX_AREA]]
- chars = [[self.word2index.get(it, self.word2index.get('<unk>')) for it in text] for text in text_list]
- if USE_API:
- requests_result = requests.post(API_URL + "/predict_product",
- json={"inputs": chars}, verify=True)
- batch_paths = json.loads(requests_result.text)['result']
- lengths = json.loads(requests_result.text)['lengths']
- else:
- lengths, scores, tran_ = sess.run([self.length, self.logit, self.tran],
- feed_dict={
- self.char_input: np.asarray(chars),
- self.dropout: 1.0
- })
- batch_paths = self.decode(scores, lengths, tran_)
- for text, path, length in zip(text_list, batch_paths, lengths):
- tags = ''.join([str(it) for it in path[:length]])
- # 提取产品
- for it in re.finditer("12*3", tags):
- start = it.start()
- end = it.end()
- if text[start:end] in ['服务', '运费']:
- continue
- _entity = Entity(doc_id=list_articles[0].id, entity_id="%s_%s_%s_%s" % (
- list_articles[0].doc_id, 0, start, end),
- entity_text=text[start:end],
- entity_type="product", sentence_index=0,
- begin_index=0, end_index=0, wordOffset_begin=start,
- wordOffset_end=end)
- list_entitys[0].append(_entity)
- product_list.append(text[start:end])
- # 提取失败原因
- for it in re.finditer("45*6", tags):
- start = it.start()
- end = it.end()
- result.append(text[start:end].replace('?', '').strip())
- reasons = []
- for it in result:
- if "(√)" in it or "(√)" in it:
- reasons = [it]
- break
- if reasons != [] and (it not in reasons[-1] and it not in reasons):
- reasons.append(it)
- elif reasons == []:
- reasons.append(it)
- if reasons == []: # 如果模型识别不到失败原因 就用规则补充
- for text in text_list:
- ser1 = re.search('\w{,4}(理由|原因):\s*((第\d+包|标项\d+|原因类型)?[::]?[\s*\w,]{2,30}((不满?足|少于|未达)((法定)?[123一二三两]家|(规定)?要求)|(项目|采购)(终止|废标)),?)+',text)
- ser2 = re.search(
- '\w{,4}(理由|原因):\s*(第\d+包|标项\d+|原因类型)?[::]?[\s*\w]{4,30},', text)
- if ser1:
- reasons.append(ser1.group(0))
- break
- elif ser2:
- reasons.append(ser2.group(0))
- break
- return {'fail_reason':';'.join(reasons)}, product_list
- if list_entitys is None:
- list_entitys = [[] for _ in range(len(list_sentences))]
- for list_sentence, list_entity in zip(list_sentences,list_entitys):
- if len(list_sentence)==0:
- result.append({"product":[]})
- continue
- # 20240827 取消,修复线上接口产品耗时长问题
- # if sentence_range: # 20240815 如果有招标内容大纲,只从前两句及大纲内提取产品,避免类似 514920213 提取错其他内容 银行流水
- # new_list = []
- # word_num = 0
- # for sentence in list_sentence:
- # if sentence.sentence_index<2:
- # new_list.append(sentence)
- # continue
- # for s1, s2 in sentence_range:
- # if sentence.sentence_index < s1:
- # continue
- # elif s1<=sentence.sentence_index <=s2:
- # new_list.append(sentence)
- # word_num += len(sentence.sentence_text)
- # elif sentence.sentence_index >= s2:
- # break
- # if word_num > 100:
- # list_sentence = new_list
- list_sentence.sort(key=lambda x:len(x.sentence_text), reverse=True)
- _begin_index = 0
- item = {"product":[]}
- temp_list = []
- while True:
- MAX_LEN = len(list_sentence[_begin_index].sentence_text)
- if MAX_LEN > MAX_AREA:
- MAX_LEN = MAX_AREA
- _LEN = MAX_AREA//MAX_LEN
- chars = [sentence.sentence_text[:MAX_LEN] for sentence in list_sentence[_begin_index:_begin_index+_LEN]]
- chars = [[self.word2index.get(it, self.word2index.get('<unk>')) for it in l] for l in chars]
- chars = pad_sequences(chars, maxlen=MAX_LEN, padding="post", truncating="post")
- if USE_API:
- requests_result = requests.post(API_URL + "/predict_product",
- json={"inputs": chars.tolist()}, verify=True)
- batch_paths = json.loads(requests_result.text)['result']
- lengths = json.loads(requests_result.text)['lengths']
- else:
- lengths, scores, tran_ = sess.run([self.length, self.logit, self.tran],
- feed_dict={
- self.char_input: np.asarray(chars),
- self.dropout: 1.0
- })
- batch_paths = self.decode(scores, lengths, tran_)
- for sentence, path, length in zip(list_sentence[_begin_index:_begin_index+_LEN],batch_paths, lengths):
- tags = ''.join([str(it) for it in path[:length]])
- for it in re.finditer("12*3", tags):
- start = it.start()
- end = it.end()
- if sentence.sentence_text[start:end] in ['服务', '运费']: # 20260522修复 739917393 去除此类型产品
- continue
- _entity = Entity(doc_id=sentence.doc_id, entity_id="%s_%s_%s_%s" % (
- sentence.doc_id, sentence.sentence_index, start, end),
- entity_text=sentence.sentence_text[start:end],
- entity_type="product", sentence_index=sentence.sentence_index,
- begin_index=0, end_index=0, wordOffset_begin=start,
- wordOffset_end=end,in_attachment=sentence.in_attachment)
- list_entity.append(_entity)
- temp_list.append(sentence.sentence_text[start:end])
- product_list.append(sentence.sentence_text[start:end])
- # item["product"] = list(set(temp_list))
- # result.append(item)
- if _begin_index+_LEN >= len(list_sentence):
- break
- _begin_index += _LEN
- item["product"] = list(set(temp_list))
- result.append(item) # 修正bug
- return {'fail_reason': ""},product_list
|