# -*- 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')) 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('')) 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