train.py 13 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300
  1. '''
  2. Created on 2019年4月22日
  3. @author: User
  4. '''
  5. # DEPRECATED(Phase1): 本文件中的 PostgreSQL 硬编码连接(host=192.168.*,
  6. # password=postgres 等)将在后续 training/ Phase 迁移到 BiddingKG.dl.infra.db。
  7. # 迁移完成前可临时使用:from BiddingKG.dl.infra.db import get_connection;
  8. # conn = get_connection("<dbname>")
  9. # 详见 ARCHITECTURE.md 第 11 章 Phase 1 与 REFACTOR_LOG.md。
  10. import sys
  11. import os
  12. sys.path.append(os.path.abspath("../../.."))
  13. from BiddingKG.dl.common.Utils import *
  14. from keras.callbacks import ModelCheckpoint
  15. from BiddingKG.dl.common.models import *
  16. import pandas as pd
  17. import keras
  18. import numpy as np
  19. os.environ["CUDA_DEVICE_ORDER"] = "PCI_BUS_ID"
  20. os.environ["CUDA_VISIBLE_DEVICES"] = ""
  21. def loadTrainData(percent=0.9,line=False):
  22. # files = ["id_token_text_begin_end_label.pk","id_token_text_begin_end_label.pk1","id_token_text_begin_end_label-selffool.pk1"]
  23. # files = ["id_token_text_begin_end_label.pk","id_token_text_begin_end_label.pk1"]
  24. files = ["id_token_text_begin_end_label-moreTrue.pk"]
  25. data_x = []
  26. data_y = []
  27. #data_id = []
  28. test_x = []
  29. test_y = []
  30. test_id = []
  31. #_,_,_,_,_,test_id_before = load("all_data_selffool.pk_line")
  32. #test_id_before = set(test_id_before)
  33. dict_label_item = dict()
  34. #统计数据分布
  35. for file in files:
  36. data = load(file)
  37. for row in data:
  38. id = row[0]
  39. label = int(row[5])
  40. if label not in dict_label_item:
  41. dict_label_item[label] = set()
  42. dict_label_item[label].add(id)
  43. dict_label_num = dict()
  44. for _key in dict_label_item.keys():
  45. dict_label_num[_key] = int(len(dict_label_item[_key])*(1-percent))
  46. for file in files:
  47. data = load(file)
  48. _count = 0
  49. for row in data:
  50. #item_x = embedding_word(spanWindow(tokens=row[1],begin_index=row[3],end_index=row[4],size=100,center_include=True,word_flag=True), shape=(3,100,60))
  51. _span = spanWindow(tokens=row[1],begin_index=row[3],end_index=row[4],size=10,center_include=True,word_flag=True,text=row[2])
  52. item_x = encodeInput(_span, word_len=50, word_flag=True,userFool=False)
  53. if line:
  54. item_x = item_x[0]+item_x[1]+item_x[2]
  55. item_y = np.zeros([6])
  56. label = int(row[5])
  57. print(_span,label)
  58. _count += 1
  59. if label not in [0,1,2,3,4,5]:
  60. continue
  61. item_y[label] = 1
  62. if np.random.random()>0.5 and dict_label_num[label]>0:
  63. dict_label_num[label] -= 1
  64. test_x.append(item_x)
  65. test_y.append(item_y)
  66. test_id.append(row[0])
  67. else:
  68. data_x.append(item_x)
  69. data_y.append(item_y)
  70. #data_id.append(row[0])
  71. # if np.random.random()>percent:
  72. # # if row[0] not in test_id_before:
  73. # data_x.append(item_x)
  74. # data_y.append(item_y)
  75. # #data_id.append(row[0])
  76. # else:
  77. # test_x.append(item_x)
  78. # test_y.append(item_y)
  79. # test_id.append(row[0])
  80. print(np.shape(np.array(data_x)),np.shape(np.array(test_x)))
  81. print(dict_label_num)
  82. if line:
  83. return np.array(data_x),np.array(data_y),np.array(test_x),np.array(test_y),None,test_id
  84. else:
  85. return np.transpose(np.array(data_x),(1,0,2)),np.array(data_y),np.transpose(np.array(test_x),(1,0,2)),np.array(test_y),None,test_id
  86. def train():
  87. # data_pk = "all_data_selffool_before-10.pk"
  88. data_pk = "all_data_selffool_moretrue-10.pk"
  89. # data_pk = "all_data_selffool_all-10.pk"
  90. if os.path.exists(data_pk):
  91. train_x,train_y,test_x,test_y,_,test_id = load(data_pk)
  92. else:
  93. train_x,train_y,test_x,test_y,_,test_id = loadTrainData()
  94. save((train_x,train_y,test_x,test_y,_,test_id),data_pk)
  95. with tf.Session(graph=tf.Graph()).as_default() as sess:
  96. with sess.graph.as_default():
  97. # dict_key_value = load("dict_key_value.pk")
  98. # model = getBiLSTMModel(input_shape=(3,50,256), vocab=fool_char_to_id.keys(), embedding_weights=dict_key_value["bert/embeddings/word_embeddings:0"], classes=6)
  99. vocab,matrix = getVocabAndMatrix(getModel_word())
  100. # model = getBiLSTMModel(input_shape=(3,50,60), vocab=vocab, embedding_weights=matrix, classes=6)
  101. model = getBiLSTMModel_entity(input_shape=(3,50,60), vocab=vocab, embedding_weights=matrix, classes=6)
  102. # model = getTextCNNModel(input_shape=(2,50,60), vocab=vocab, embedding_weights=matrix, classes=6)
  103. '''
  104. for k,v in dict_key_value.items():
  105. if re.search("encoder",k) is not None:
  106. sess.run(tf.assign(sess.graph.get_tensor_by_name(k[13:]),v))
  107. print(k)
  108. '''
  109. #model = getTextCNNModel(input_shape=(3,50,60), vocab=vocab, embedding_weights=weights, classes=6)
  110. # model.load_weights("log/ep044-loss0.142-val_loss0.200-f1_score0.934.h5",skip_mismatch=True,by_name=True)
  111. model.load_weights("log/min_val_loss_ep027-loss0.112-val_loss0.109-f1_score0.963.h5")
  112. #model.summary()
  113. #print("11111111111",sess.run(sess.graph.get_tensor_by_name("encoder/layer_0/attention/self/query/kernel:0")))
  114. callback = ModelCheckpoint(filepath="log/"+"min_val_loss_ep{epoch:03d}-loss{loss:.3f}-val_loss{val_loss:.3f}-f1_score{val_f1_score:.3f}.h5",monitor="val_loss",save_best_only=True, save_weights_only=True, mode="min")
  115. callback1 = ModelCheckpoint(filepath="log/"+"min_loss_ep{epoch:03d}-loss{loss:.3f}-val_loss{val_loss:.3f}-f1_score{val_f1_score:.3f}.h5",monitor="loss",save_best_only=True, save_weights_only=True, mode="min")
  116. history_model = model.fit(x=[train_x[0],train_x[1],train_x[2]],y=train_y,validation_data=([test_x[0],test_x[1],test_x[2]],test_y),epochs=600,batch_size=96,shuffle=True,callbacks=[callback,callback1])
  117. # history_model = model.fit(x=[train_x[0],train_x[1]],y=train_y,validation_data=([test_x[0],test_x[1]],test_y),epochs=600,batch_size=128,shuffle=True,callbacks=[callback,callback1])
  118. # history_model = model.fit(x=train_x,y=train_y,validation_data=(test_x,test_y),epochs=600,batch_size=128,shuffle=True,callbacks=[callback])
  119. #print("2222222222222",sess.run(sess.graph.get_tensor_by_name("encoder/layer_0/attention/self/query/kernel:0")))
  120. def test():
  121. _span = [':预算金额1000000元,中标金额', '1df元', ';']
  122. _input = encodeInput(_span, word_len=50, word_flag=True,userFool=True)
  123. print(_input)
  124. print(len(_input))
  125. print(len(_input[0]))
  126. print(len(_input[1]))
  127. print(len(_input[2]))
  128. def statis():
  129. df = pd.read_excel("测试数据_role-biws-biw0.xls")
  130. result = {"正确-词":0,
  131. "错误-词":0,
  132. "正确-字":0,
  133. "错误-字":0}
  134. for i in range(6):
  135. result["正确-词"+str(i)] = 0
  136. result["错误-词"+str(i)] = 0
  137. result["正确-字"+str(i)] = 0
  138. result["错误-字"+str(i)] = 0
  139. for label_ws,prob_ws,label_w,prob_w,label_true in zip(df["list_newlabel"],df["list_newprob"],df["list_newlabel_cnn"],df["list_newprob_cnn"],df["label_true"]):
  140. if int(label_ws)==int(label_true):
  141. key = "正确-词"
  142. result[key] += 1
  143. result[key+str(int(label_ws))]+=1
  144. else:
  145. key = "错误-词"
  146. result[key] += 1
  147. result[key+str(int(label_ws))]+=1
  148. if int(label_w)==int(label_true):
  149. key = "正确-字"
  150. result[key] += 1
  151. result[key+str(int(label_w))]+=1
  152. else:
  153. key = "错误-字"
  154. result[key] += 1
  155. result[key+str(int(label_w))]+=1
  156. data = []
  157. for key in result.keys():
  158. data.append([key,result[key]])
  159. data.sort(key=lambda x:x[0])
  160. for item in data:
  161. print(item)
  162. def val():
  163. data_pk = "all_data_selffool.pk_line"
  164. train_x,train_y,test_x,test_y,_,test_id = load(data_pk)
  165. vocab,matrix = getVocabAndMatrix(getModel_word())
  166. model = getBiLSTMModel(input_shape=(1,150,60), vocab=vocab, embedding_weights=matrix, classes=6)
  167. model.load_weights("log/ep064-loss0.585-val_loss0.634-f1_score0.927.h5")
  168. # predict_y = np.argmax(model.predict([test_x[0],test_x[1],test_x[2]]),-1)
  169. predict_y = np.argmax(model.predict(test_x),-1)
  170. dict_notTrue = dict()
  171. for _y,Y,_id in zip(predict_y,np.argmax(test_y,-1),test_id):
  172. if _y!=Y:
  173. dict_notTrue[_id] = [_y,Y]
  174. token_data = load("id_token_text_begin_end_label-selffool.pk1")
  175. test_before = []
  176. test_center = []
  177. test_after = []
  178. test_label = []
  179. test_predict = []
  180. for item in token_data:
  181. if item[0] in dict_notTrue:
  182. token = item[1]
  183. text = item[2]
  184. begin = item[3]
  185. end = item[4]
  186. predict,label = dict_notTrue[item[0]]
  187. _span = spanWindow(tokens=token,begin_index=begin,end_index=end,size=10,center_include=True,word_flag=True,text=text)
  188. before,center,after = _span
  189. test_before.append(before)
  190. test_center.append(center)
  191. test_after.append(after)
  192. test_label.append(label)
  193. test_predict.append(predict)
  194. data = {"test_before":test_before,"test_center":test_center,"test_after":test_after,"test_label":test_label,"test_predict":test_predict}
  195. df = pd.DataFrame(data)
  196. df.to_excel("val_bert_position.xls",columns=["test_before","test_center","test_after","test_label","test_predict"])
  197. def get_savedmodel():
  198. with tf.Session(graph=tf.Graph()).as_default() as sess:
  199. with sess.graph.as_default():
  200. vocab,matrix = getVocabAndMatrix(getModel_word(),Embedding_size=60)
  201. # model = getBiLSTMModel(input_shape=(3,50,60), vocab=vocab, embedding_weights=matrix, classes=6)
  202. model = getBiLSTMModel_entity(input_shape=(3,50,60), vocab=vocab, embedding_weights=matrix, classes=6)
  203. # model = getTextCNNModel(input_shape=(2,50,60), vocab=vocab, embedding_weights=matrix, classes=6)
  204. # filepath = "log/ep001-loss0.087-val_loss0.172-f1_score0.944.h5"
  205. filepath = "../../dl_dev/role/log/min_val_loss_ep034-loss0.070-val_loss0.068-f1_score0.975.h5"
  206. model.load_weights(filepath)
  207. tf.saved_model.simple_save(sess,
  208. "role_savedmodel/",
  209. inputs={"input0":model.input[0],
  210. "input1":model.input[1],
  211. "input2":model.input[2]},
  212. outputs={"outputs":model.output}
  213. )
  214. def get_tensorboard():
  215. with tf.Session(graph=tf.Graph()) as sess:
  216. tf.saved_model.loader.load(sess,export_dir="role_savedmodel",tags=["serve"])
  217. writer = tf.summary.FileWriter(graph=sess.graph,logdir="log2")
  218. def relabel():
  219. list_id = []
  220. list_before = []
  221. list_center = []
  222. list_after = []
  223. list_label = []
  224. files = ["id_token_text_begin_end_label.pk","id_token_text_begin_end_label.pk1"]
  225. for file in files:
  226. data = load(file)
  227. _count = 0
  228. for row in data:
  229. #item_x = embedding_word(spanWindow(tokens=row[1],begin_index=row[3],end_index=row[4],size=100,center_include=True,word_flag=True), shape=(3,100,60))
  230. _span = spanWindow(tokens=row[1],begin_index=row[3],end_index=row[4],size=15,center_include=True,word_flag=True,text=row[2])
  231. _label = row[5]
  232. if int(_label) in [3,4]:
  233. list_id.append(row[0])
  234. list_before.append(_span[0])
  235. list_center.append(_span[1])
  236. list_after.append(_span[2])
  237. list_label.append(str(_label))
  238. df = pd.DataFrame({"list_id":list_id,
  239. "list_before":list_before,
  240. "list_center":list_center,
  241. "list_after":list_after,
  242. "list_label":list_label})
  243. df.to_excel("relabel_1.xls",columns=["list_id","list_before","list_center","list_after","list_label"])
  244. def generate_data():
  245. file_before = "D:\\myProject\\traindata\\"
  246. files = ["id_token_text_begin_end_label.pk","id_token_text_begin_end_label.pk1"]
  247. data = load(file_before+"id_token_text_begin_end_label-selffool.pk1")
  248. df = pd.read_excel(file_before+"relabel_1.xls")
  249. set_id = set(df["list_id"])
  250. for file in files:
  251. temp_data = load(file_before+file)
  252. for row in temp_data:
  253. if row[0] in set_id:
  254. # print(row)
  255. data.append(row)
  256. save(data,file_before+"id_token_text_begin_end_label-moreTrue.pk")
  257. if __name__=="__main__":
  258. # loadTrainData()
  259. train()
  260. # relabel()
  261. # generate_data()
  262. test()
  263. #statis()
  264. # val()
  265. # get_savedmodel()
  266. # get_tensorboard()
  267. pass