time.py 5.5 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129
  1. # -*- coding: utf-8 -*-
  2. """``TimePredictor`` — 时间分类模型。
  3. Phase 5 从 ``interface/predictor.py``(约 3028-3132 行)迁出。
  4. 原 ``from common.Utils import *`` / ``from interface.modelFactory 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 numpy as np
  12. import tensorflow as tf
  13. from BiddingKG.dl.common.logging import log
  14. from BiddingKG.dl.common.Utils import timeFormat
  15. from BiddingKG.dl.common.context_utils import spanWindow
  16. from BiddingKG.dl.model_runtime.embed import getModel_w2v
  17. from BiddingKG.dl.predictors._common import INTERFACE_DIR
  18. from BiddingKG.dl.services.external_inference.paieas import limitRun
  19. __all__ = ["TimePredictor"]
  20. class TimePredictor():
  21. def __init__(self,config=None):
  22. self.sess = tf.Session(graph=tf.Graph(),config=config)
  23. self.inputs_code = None
  24. self.outputs_code = None
  25. self.input_shape = (2,40,128)
  26. self.load_model()
  27. def load_model(self):
  28. model_path = os.path.join(INTERFACE_DIR, 'timesplit_model')
  29. if self.inputs_code is None:
  30. log("get model of time")
  31. with self.sess.as_default():
  32. with self.sess.graph.as_default():
  33. meta_graph_def = tf.saved_model.loader.load(self.sess, tags=["serve"], export_dir=model_path)
  34. signature_key = tf.saved_model.signature_constants.DEFAULT_SERVING_SIGNATURE_DEF_KEY
  35. signature_def = meta_graph_def.signature_def
  36. self.inputs_code = []
  37. self.inputs_code.append(
  38. self.sess.graph.get_tensor_by_name(signature_def[signature_key].inputs["input0"].name))
  39. self.inputs_code.append(
  40. self.sess.graph.get_tensor_by_name(signature_def[signature_key].inputs["input1"].name))
  41. self.outputs_code = self.sess.graph.get_tensor_by_name(signature_def[signature_key].outputs["outputs"].name)
  42. return self.inputs_code, self.outputs_code
  43. else:
  44. return self.inputs_code, self.outputs_code
  45. def search_time_data(self,list_sentences,list_entitys):
  46. data_x = []
  47. points_entitys = []
  48. for list_sentence, list_entity in zip(list_sentences, list_entitys):
  49. p_entitys = 0
  50. p_sentences = 0
  51. list_sentence.sort(key=lambda x: x.sentence_index)
  52. while(p_entitys<len(list_entity)):
  53. entity = list_entity[p_entitys]
  54. if entity.entity_type in ['time']:
  55. while(p_sentences<len(list_sentence)):
  56. sentence = list_sentence[p_sentences]
  57. if entity.doc_id == sentence.doc_id and entity.sentence_index == sentence.sentence_index:
  58. # left = sentence.sentence_text[max(0,entity.wordOffset_begin-self.input_shape[1]):entity.wordOffset_begin]
  59. # right = sentence.sentence_text[entity.wordOffset_end:entity.wordOffset_end+self.input_shape[1]]
  60. s = spanWindow(tokens=sentence.tokens,begin_index=entity.begin_index,end_index=entity.end_index,size=self.input_shape[1])
  61. left = s[0]
  62. right = s[1]
  63. context = [left, right]
  64. x = self.embedding_words(context, shape=self.input_shape)
  65. data_x.append(x)
  66. points_entitys.append(entity)
  67. break
  68. p_sentences += 1
  69. p_entitys += 1
  70. if len(points_entitys)==0:
  71. return None
  72. data_x = np.transpose(np.array(data_x), (1, 0, 2, 3))
  73. return [data_x, points_entitys]
  74. def embedding_words(self, datas, shape):
  75. '''
  76. @summary:查找词汇对应的词向量
  77. @param:
  78. datas:词汇的list
  79. shape:结果的shape
  80. @return: array,返回对应shape的词嵌入
  81. '''
  82. model_w2v = getModel_w2v()
  83. embed = np.zeros(shape)
  84. length = shape[1]
  85. out_index = 0
  86. for data in datas:
  87. index = 0
  88. for item in data:
  89. item_not_space = re.sub("\s*", "", item)
  90. if index >= length:
  91. break
  92. if item_not_space in model_w2v.vocab:
  93. embed[out_index][index] = model_w2v[item_not_space]
  94. index += 1
  95. else:
  96. embed[out_index][index] = model_w2v['unk']
  97. index += 1
  98. out_index += 1
  99. return embed
  100. def predict(self, list_sentences,list_entitys):
  101. datas = self.search_time_data(list_sentences, list_entitys)
  102. if datas is None:
  103. return
  104. points_entitys = datas[1]
  105. with self.sess.as_default():
  106. predict_y = limitRun(self.sess,[self.outputs_code], feed_dict={self.inputs_code[0]:datas[0][0]
  107. ,self.inputs_code[1]:datas[0][1]})[0]
  108. for i in range(len(predict_y)):
  109. entity = points_entitys[i]
  110. label = np.argmax(predict_y[i])
  111. values = []
  112. for item in predict_y[i]:
  113. values.append(item)
  114. if label != 0:
  115. if not timeFormat(entity.entity_text):
  116. label = 0
  117. values[0] = 0.5
  118. entity.set_Role(label, values)