# -*- coding: utf-8 -*- """``TimePredictor`` — 时间分类模型。 Phase 5 从 ``interface/predictor.py``(约 3028-3132 行)迁出。 原 ``from common.Utils import *`` / ``from interface.modelFactory import *`` 已替换为显式 import;``os.path.dirname(__file__)`` 路径引用替换为 ``predictors._common.INTERFACE_DIR``。 """ from __future__ import absolute_import import os import re import numpy as np import tensorflow as tf from BiddingKG.dl.common.logging import log from BiddingKG.dl.common.Utils import timeFormat from BiddingKG.dl.common.context_utils import spanWindow from BiddingKG.dl.model_runtime.embed import getModel_w2v from BiddingKG.dl.predictors._common import INTERFACE_DIR from BiddingKG.dl.services.external_inference.paieas import limitRun __all__ = ["TimePredictor"] class TimePredictor(): def __init__(self,config=None): self.sess = tf.Session(graph=tf.Graph(),config=config) self.inputs_code = None self.outputs_code = None self.input_shape = (2,40,128) self.load_model() def load_model(self): model_path = os.path.join(INTERFACE_DIR, 'timesplit_model') if self.inputs_code is None: log("get model of time") with self.sess.as_default(): with self.sess.graph.as_default(): meta_graph_def = tf.saved_model.loader.load(self.sess, tags=["serve"], export_dir=model_path) signature_key = tf.saved_model.signature_constants.DEFAULT_SERVING_SIGNATURE_DEF_KEY signature_def = meta_graph_def.signature_def self.inputs_code = [] self.inputs_code.append( self.sess.graph.get_tensor_by_name(signature_def[signature_key].inputs["input0"].name)) self.inputs_code.append( self.sess.graph.get_tensor_by_name(signature_def[signature_key].inputs["input1"].name)) self.outputs_code = self.sess.graph.get_tensor_by_name(signature_def[signature_key].outputs["outputs"].name) return self.inputs_code, self.outputs_code else: return self.inputs_code, self.outputs_code def search_time_data(self,list_sentences,list_entitys): data_x = [] points_entitys = [] for list_sentence, list_entity in zip(list_sentences, list_entitys): p_entitys = 0 p_sentences = 0 list_sentence.sort(key=lambda x: x.sentence_index) while(p_entitys= length: break if item_not_space in model_w2v.vocab: embed[out_index][index] = model_w2v[item_not_space] index += 1 else: embed[out_index][index] = model_w2v['unk'] index += 1 out_index += 1 return embed def predict(self, list_sentences,list_entitys): datas = self.search_time_data(list_sentences, list_entitys) if datas is None: return points_entitys = datas[1] with self.sess.as_default(): predict_y = limitRun(self.sess,[self.outputs_code], feed_dict={self.inputs_code[0]:datas[0][0] ,self.inputs_code[1]:datas[0][1]})[0] for i in range(len(predict_y)): entity = points_entitys[i] label = np.argmax(predict_y[i]) values = [] for item in predict_y[i]: values.append(item) if label != 0: if not timeFormat(entity.entity_text): label = 0 values[0] = 0.5 entity.set_Role(label, values)