save_models.py 13 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302303304305306307308309310311312313314315316317318319320321322323324325326327328329330331332333334335336337
  1. # -*- coding: utf-8 -*-
  2. """Training / model-saving helpers.
  3. 按 ARCHITECTURE.md Phase 5 拆分建议,从 ``interface/predictor.py`` 迁出的
  4. 训练 / 模型保存相关函数。这些函数构建或加载 Keras / TensorFlow 模型,并将其
  5. 导出为 TensorFlow ``SavedModel`` 产物以供部署调用(codename / role / money /
  6. person / form / codesplit / timesplit)。
  7. 原位置:``interface/predictor.py`` 末尾(约第 10155-10448 行)。
  8. """
  9. from __future__ import absolute_import
  10. import os
  11. import re
  12. import sys
  13. import pickle
  14. import h5py
  15. import numpy as np
  16. import tensorflow as tf
  17. from keras import models, layers
  18. # from keras_contrib.layers.crf import CRF
  19. from BiddingKG.dl.common.Utils import load, precision, recall, f1_score
  20. from BiddingKG.dl.complaint.punishNo_tf import BiLSTM_CRF_tfmodel
  21. from BiddingKG.dl.predictors.prem import PREMPredict, EPCPredict
  22. from BiddingKG.dl.predictors.form import FormPredictor
  23. from BiddingKG.dl.predictors._common import INTERFACE_DIR
  24. __all__ = [
  25. "getSavedModel",
  26. "getBiLSTMCRFModel",
  27. "h5_to_graph",
  28. "initialize_uninitialized",
  29. "save_codename_model",
  30. "save_role_model",
  31. "save_money_model",
  32. "save_person_model",
  33. "save_form_model",
  34. "save_codesplit_model",
  35. "save_timesplit_model",
  36. ]
  37. def getSavedModel():
  38. #predictor = FormPredictor()
  39. graph = tf.Graph()
  40. with graph.as_default():
  41. model = tf.keras.models.load_model("../form/model/model_form.model_item.hdf5",custom_objects={"precision":precision,"recall":recall,"f1_score":f1_score})
  42. #print(tf.graph_util.remove_training_nodes(model))
  43. tf.saved_model.simple_save(
  44. tf.keras.backend.get_session(),
  45. "./h5_savedmodel/",
  46. inputs={"image": model.input},
  47. outputs={"scores": model.output}
  48. )
  49. def getBiLSTMCRFModel(MAX_LEN,vocab,EMBED_DIM,BiRNN_UNITS,chunk_tags,weights):
  50. '''
  51. model = models.Sequential()
  52. model.add(layers.Embedding(len(vocab), EMBED_DIM, mask_zero=True)) # Random embedding
  53. model.add(layers.Bidirectional(layers.LSTM(BiRNN_UNITS // 2, return_sequences=True)))
  54. crf = CRF(len(chunk_tags), sparse_target=True)
  55. model.add(crf)
  56. model.summary()
  57. model.compile('adam', loss=crf.loss_function, metrics=[crf.accuracy])
  58. return model
  59. '''
  60. input = layers.Input(shape=(None,),dtype="int32")
  61. if weights is not None:
  62. embedding = layers.embeddings.Embedding(len(vocab),EMBED_DIM,mask_zero=True,weights=[weights],trainable=True)(input)
  63. else:
  64. embedding = layers.embeddings.Embedding(len(vocab),EMBED_DIM,mask_zero=True)(input)
  65. bilstm = layers.Bidirectional(layers.LSTM(BiRNN_UNITS//2,return_sequences=True))(embedding)
  66. bilstm_dense = layers.TimeDistributed(layers.Dense(len(chunk_tags)))(bilstm)
  67. crf = CRF(len(chunk_tags),sparse_target=True)
  68. crf_out = crf(bilstm_dense)
  69. model = models.Model(input=[input],output = [crf_out])
  70. model.summary()
  71. model.compile(optimizer = 'adam', loss = crf.loss_function, metrics = [crf.accuracy])
  72. return model
  73. def h5_to_graph(sess,graph,h5file):
  74. f = h5py.File(h5file,'r') #打开h5文件
  75. def getValue(v):
  76. _value = f["model_weights"]
  77. list_names = str(v.name).split("/")
  78. for _index in range(len(list_names)):
  79. print(v.name)
  80. if _index==1:
  81. _value = _value[list_names[0]]
  82. _value = _value[list_names[_index]]
  83. return _value
  84. def _load_attributes_from_hdf5_group(group, name):
  85. """Loads attributes of the specified name from the HDF5 group.
  86. This method deals with an inherent problem
  87. of HDF5 file which is not able to store
  88. data larger than HDF5_OBJECT_HEADER_LIMIT bytes.
  89. # Arguments
  90. group: A pointer to a HDF5 group.
  91. name: A name of the attributes to load.
  92. # Returns
  93. data: Attributes data.
  94. """
  95. if name in group.attrs:
  96. data = [n.decode('utf8') for n in group.attrs[name]]
  97. else:
  98. data = []
  99. chunk_id = 0
  100. while ('%s%d' % (name, chunk_id)) in group.attrs:
  101. data.extend([n.decode('utf8')
  102. for n in group.attrs['%s%d' % (name, chunk_id)]])
  103. chunk_id += 1
  104. return data
  105. def readGroup(gr,parent_name,data):
  106. for subkey in gr:
  107. print(subkey)
  108. if parent_name!=subkey:
  109. if parent_name=="":
  110. _name = subkey
  111. else:
  112. _name = parent_name+"/"+subkey
  113. else:
  114. _name = parent_name
  115. if str(type(gr[subkey]))=="<class 'h5py._hl.group.Group'>":
  116. readGroup(gr[subkey],_name,data)
  117. else:
  118. data.append([_name,gr[subkey].value])
  119. print(_name,gr[subkey].shape)
  120. layer_names = _load_attributes_from_hdf5_group(f["model_weights"], 'layer_names')
  121. list_name_value = []
  122. readGroup(f["model_weights"], "", list_name_value)
  123. '''
  124. for k, name in enumerate(layer_names):
  125. g = f["model_weights"][name]
  126. weight_names = _load_attributes_from_hdf5_group(g, 'weight_names')
  127. #weight_values = [np.asarray(g[weight_name]) for weight_name in weight_names]
  128. for weight_name in weight_names:
  129. list_name_value.append([weight_name,np.asarray(g[weight_name])])
  130. '''
  131. for name_value in list_name_value:
  132. name = name_value[0]
  133. '''
  134. if re.search("dense",name) is not None:
  135. name = name[:7]+"_1"+name[7:]
  136. '''
  137. value = name_value[1]
  138. print(name,graph.get_tensor_by_name(name),np.shape(value))
  139. sess.run(tf.assign(graph.get_tensor_by_name(name),value))
  140. def initialize_uninitialized(sess):
  141. global_vars = tf.global_variables()
  142. is_not_initialized = sess.run([tf.is_variable_initialized(var) for var in global_vars])
  143. not_initialized_vars = [v for (v, f) in zip(global_vars, is_not_initialized) if not f]
  144. adam_vars = []
  145. for _vars in not_initialized_vars:
  146. if re.search("Adam",_vars.name) is not None:
  147. adam_vars.append(_vars)
  148. print([str(i.name) for i in adam_vars]) # only for testing
  149. if len(adam_vars):
  150. sess.run(tf.variables_initializer(adam_vars))
  151. def save_codename_model():
  152. # filepath = "../projectCode/models/model_project_"+str(60)+"_"+str(200)+".hdf5"
  153. filepath = "../../dl_dev/projectCode/models_tf/59-L0.471516189943-F0.8802154826344823-P0.8789179683459191-R0.8815168335321886/model.ckpt"
  154. vocabpath = "../projectCode/models/vocab.pk"
  155. classlabelspath = "../projectCode/models/classlabels.pk"
  156. # vocab = load(vocabpath)
  157. # class_labels = load(classlabelspath)
  158. w2v_matrix = load('codename_w2v_matrix.pk')
  159. graph = tf.get_default_graph()
  160. with graph.as_default() as g:
  161. ''''''
  162. # model = getBiLSTMCRFModel(None, vocab, 60, 200, class_labels,weights=None)
  163. #model = models.load_model(filepath,custom_objects={'precision':precision,'recall':recall,'f1_score':f1_score,"CRF":CRF,"loss":CRF.loss_function})
  164. sess = tf.Session(graph=g)
  165. # sess = tf.keras.backend.get_session()
  166. char_input, logits, target, keepprob, length, crf_loss, trans, train_op = BiLSTM_CRF_tfmodel(sess, w2v_matrix)
  167. #with sess.as_default():
  168. sess.run(tf.global_variables_initializer())
  169. # print(sess.run("time_distributed_1/kernel:0"))
  170. # model.load_weights(filepath)
  171. saver = tf.train.Saver()
  172. saver.restore(sess, filepath)
  173. # print("logits",sess.run(logits))
  174. # print("#",sess.run("time_distributed_1/kernel:0"))
  175. # x = load("codename_x.pk")
  176. #y = model.predict(x)
  177. # y = sess.run(model.output,feed_dict={model.input:x})
  178. # for item in np.argmax(y,-1):
  179. # print(item)
  180. tf.saved_model.simple_save(
  181. sess,
  182. "./codename_savedmodel_tf/",
  183. inputs={"inputs": char_input,
  184. "inputs_length":length,
  185. 'keepprob':keepprob},
  186. outputs={"logits": logits,
  187. "trans":trans}
  188. )
  189. def save_role_model():
  190. '''
  191. @summary: 保存model为savedModel,部署到PAI平台上调用
  192. '''
  193. model_role = PREMPredict().model_role
  194. with model_role.graph.as_default():
  195. model = model_role.getModel()
  196. sess = tf.Session(graph=model_role.graph)
  197. print(type(model.input))
  198. sess.run(tf.global_variables_initializer())
  199. h5_to_graph(sess, model_role.graph, model_role.model_role_file)
  200. model = model_role.getModel()
  201. tf.saved_model.simple_save(sess,
  202. "./role_savedmodel/",
  203. inputs={"input0":model.input[0],
  204. "input1":model.input[1],
  205. "input2":model.input[2]},
  206. outputs={"outputs":model.output}
  207. )
  208. def save_money_model():
  209. model_file = os.path.join(INTERFACE_DIR, "../money/models/model_money_word.h5")
  210. graph = tf.Graph()
  211. with graph.as_default():
  212. sess = tf.Session(graph=graph)
  213. with sess.as_default():
  214. # model = model_money.getModel()
  215. # model.summary()
  216. # sess.run(tf.global_variables_initializer())
  217. # h5_to_graph(sess, model_money.graph, model_money.model_money_file)
  218. model = models.load_model(model_file,custom_objects={'precision':precision,'recall':recall,'f1_score':f1_score})
  219. model.summary()
  220. print(model.weights)
  221. tf.saved_model.simple_save(sess,
  222. "./money_savedmodel2/",
  223. inputs = {"input0":model.input[0],
  224. "input1":model.input[1],
  225. "input2":model.input[2]},
  226. outputs = {"outputs":model.output}
  227. )
  228. def save_person_model():
  229. model_person = EPCPredict().model_person
  230. with model_person.graph.as_default():
  231. x = load("person_x.pk")
  232. _data = np.transpose(np.array(x),(1,0,2,3))
  233. model = model_person.getModel()
  234. sess = tf.Session(graph=model_person.graph)
  235. with sess.as_default():
  236. sess.run(tf.global_variables_initializer())
  237. model_person.load_weights()
  238. #h5_to_graph(sess, model_person.graph, model_person.model_person_file)
  239. predict_y = sess.run(model.output,feed_dict={model.input[0]:_data[0],model.input[1]:_data[1]})
  240. #predict_y = model.predict([_data[0],_data[1]])
  241. print(np.argmax(predict_y,-1))
  242. tf.saved_model.simple_save(sess,
  243. "./person_savedmodel/",
  244. inputs={"input0":model.input[0],
  245. "input1":model.input[1]},
  246. outputs = {"outputs":model.output})
  247. def save_form_model():
  248. model_form = FormPredictor()
  249. with model_form.graph.as_default():
  250. model = model_form.getModel("item")
  251. sess = tf.Session(graph=model_form.graph)
  252. sess.run(tf.global_variables_initializer())
  253. h5_to_graph(sess, model_form.graph, model_form.model_file_item)
  254. tf.saved_model.simple_save(sess,
  255. "./form_savedmodel/",
  256. inputs={"inputs":model.input},
  257. outputs = {"outputs":model.output})
  258. def save_codesplit_model():
  259. filepath_code = "../../dl_dev/projectCode/models/model_code.hdf5"
  260. graph = tf.Graph()
  261. with graph.as_default():
  262. model_code = models.load_model(filepath_code, custom_objects={'precision':precision,'recall':recall,'f1_score':f1_score})
  263. sess = tf.Session()
  264. sess.run(tf.global_variables_initializer())
  265. h5_to_graph(sess, graph, filepath_code)
  266. tf.saved_model.simple_save(sess,
  267. "./codesplit_savedmodel/",
  268. inputs={"input0":model_code.input[0],
  269. "input1":model_code.input[1],
  270. "input2":model_code.input[2]},
  271. outputs={"outputs":model_code.output})
  272. def save_timesplit_model():
  273. filepath = '../time/model_label_time_classify.model.hdf5'
  274. with tf.Graph().as_default() as graph:
  275. time_model = models.load_model(filepath, custom_objects={'precision': precision, 'recall': recall, 'f1_score': f1_score})
  276. with tf.Session() as sess:
  277. sess.run(tf.global_variables_initializer())
  278. h5_to_graph(sess, graph, filepath)
  279. tf.saved_model.simple_save(sess,
  280. "./timesplit_model/",
  281. inputs={"input0":time_model.input[0],
  282. "input1":time_model.input[1]},
  283. outputs={"outputs":time_model.output})