train.py 11 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244
  1. #from general_data import getTokensLabels
  2. # DEPRECATED(Phase1): 本文件中的 PostgreSQL 硬编码连接(host=192.168.*,
  3. # password=postgres 等)将在后续 training/ Phase 迁移到 BiddingKG.dl.infra.db。
  4. # 迁移完成前可临时使用:from BiddingKG.dl.infra.db import get_connection;
  5. # conn = get_connection("<dbname>")
  6. # 详见 ARCHITECTURE.md 第 11 章 Phase 1 与 REFACTOR_LOG.md。
  7. import sys
  8. import os
  9. sys.path.append(os.path.abspath("../.."))
  10. # from model import *
  11. from keras.callbacks import ModelCheckpoint
  12. from keras import layers,models,optimizers,losses
  13. import psycopg2
  14. from BiddingKG.dl.common.Utils import *
  15. import pandas as pd
  16. from BiddingKG.dl.common.models import *
  17. # os.environ["CUDA_DEVICE_ORDER"] = "PCI_BUS_ID"
  18. # os.environ["CUDA_VISIBLE_DEVICES"] = ""
  19. sourcetable = "label_guest_person"
  20. domain = sourcetable.split("_")[2]
  21. model_file = "model_"+domain+".model"
  22. input_shape = (2,10,128)
  23. output_shape = [5]
  24. def getTokensLabels(t,isTrain=True):
  25. '''
  26. @summary: 取得模型的输入输出数据
  27. @param:
  28. t:标签数据所在表
  29. @return: type:array,array,list meaning:输入,输出,实体id
  30. '''
  31. conn = psycopg2.connect(dbname="BiddingKG",user="postgres",password="postgres",host="192.168.2.101")
  32. cursor = conn.cursor()
  33. if isTrain:
  34. sql = " select B.tokens,A.begin_index,A.end_index,C.label,A.entity_id from train_entity_copy A,train_sentences_copy B,"+t+" C where A.doc_id=B.doc_id and A.sentence_index=B.sentence_index and A.entity_type='person' and A.entity_id=C.entity_id and C.entity_id not in (select entity_id from "+t+" order by entity_id limit 2000)"
  35. else:
  36. sql = " select B.tokens,A.begin_index,A.end_index,C.label,A.entity_id from train_entity_copy A,train_sentences_copy B,"+t+" C where A.doc_id=B.doc_id and A.sentence_index=B.sentence_index and A.entity_type='person' and A.entity_id=C.entity_id and C.entity_id in (select entity_id from "+t+" order by entity_id limit 2000)"
  37. cursor.execute(sql)
  38. print(sql)
  39. data_x = []
  40. data_y = []
  41. data_context = []
  42. rows = cursor.fetchmany(1000)
  43. allLimit = 250000
  44. all = 0
  45. i = 0
  46. while(rows):
  47. for row in rows:
  48. if all>=allLimit:
  49. break
  50. item_x = embedding(spanWindow(tokens=row[0],begin_index=row[1],end_index=row[2],size=input_shape[1]),shape=input_shape)
  51. # item_x = encodeInput(spanWindow(tokens=row[0],begin_index=row[1],end_index=row[2],size=10), word_len=50, word_flag=True,userFool=False)
  52. # _span = spanWindow(tokens=row[0],begin_index=row[1],end_index=row[2],size=10,word_flag=False)
  53. # item_x = encodeInput(_span, word_len=10, word_flag=False,userFool=False)
  54. item_y = np.zeros(output_shape)
  55. item_y[row[3]] = 1
  56. all += 1
  57. if not isTrain:
  58. item_context = []
  59. item_context.append(row[4])
  60. data_context.append(item_context)
  61. data_x.append(item_x)
  62. data_y.append(item_y)
  63. i += 1
  64. rows = cursor.fetchmany(1000)
  65. return np.transpose(np.array(data_x),(1,0,2,3)),np.array(data_y),data_context
  66. def getBiRNNModel():
  67. '''
  68. @summary: 获得模型
  69. '''
  70. L_input = layers.Input(shape=input_shape[1:],dtype="float32")
  71. #C_input = layers.Input(shape=(10,128),dtype="float32")
  72. R_input = layers.Input(shape=input_shape[1:],dtype="float32")
  73. #lstm_0 = layers.Bidirectional(layers.LSTM(16,return_sequences=True))(ThreeBilstm(0)(input))
  74. lstm_0 = layers.Bidirectional(layers.LSTM(32,return_sequences=True))(L_input)
  75. avg_0 = layers.GlobalAveragePooling1D()(lstm_0)
  76. #lstm_1 = layers.Bidirectional(layers.LSTM(16,return_sequences=True))(C_input)
  77. #avg_1 = layers.GlobalAveragePooling1D()(lstm_1)
  78. lstm_2 = layers.Bidirectional(layers.LSTM(32,return_sequences=True))(R_input)
  79. avg_2 = layers.GlobalAveragePooling1D()(lstm_2)
  80. #concat = layers.merge([avg_0,avg_1,avg_2],mode="concat")
  81. concat = layers.merge([avg_0,avg_2],mode="concat")
  82. output = layers.Dense(output_shape[0],activation="softmax")(concat)
  83. model = models.Model(inputs=[L_input,R_input],outputs=output)
  84. model.compile(optimizer=optimizers.Adam(lr=0.0005),loss=losses.binary_crossentropy,metrics=[precision,recall,f1_score])
  85. return model
  86. def training():
  87. '''
  88. @summary: 训练模型
  89. '''
  90. model = getBiRNNModel()
  91. model.summary()
  92. train_x,train_y,_ = getTokensLabels(isTrain=True,t="hand_label_person")
  93. #print(np.shape(train_x))
  94. test_x,test_y,test_context = getTokensLabels(isTrain=False,t="hand_label_person")
  95. checkpoint = ModelCheckpoint(model_file+".hdf5",monitor="val_loss",verbose=1,save_best_only=True,mode='min')
  96. 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=100,batch_size=256,shuffle=True,callbacks=[checkpoint])
  97. # predict_y = model.predict([test_x[0],test_x[1]])
  98. #
  99. # conn = psycopg2.connect(dbname='BiddingKG', user='postgres',password='postgres',host='192.168.2.101')
  100. # cursor = conn.cursor()
  101. # table = 'predict_person'
  102. # cursor.execute(" select to_regclass('"+table+"') is null ")
  103. # notExists = cursor.fetchall()[0][0]
  104. # if notExists:
  105. # cursor.execute(" create table "+table+" (entity_id text,predect int,label int)")
  106. # else:
  107. # cursor.execute(" delete from "+table)
  108. #
  109. #
  110. #
  111. # with open("predict.txt","w",encoding="utf8") as f:
  112. # for i in range(len(predict_y)):
  113. # if np.argmax(predict_y[i]) != np.argmax(test_y[i]):
  114. # f.write("\n")
  115. # f.write(str(test_context[i][0]))
  116. # f.write("\t")
  117. # f.write(str(np.argmax(predict_y[i])))
  118. # f.write("\t")
  119. # f.write(str(np.argmax(test_y[i])))
  120. # f.write("\n")
  121. # sql = " insert into "+table+"(entity_id ,predect ,label) values('"+str(test_context[i][0])+"','"+str(int(np.argmax(predict_y[i])))+"','"+str(int(np.argmax(test_y[i])))+"')"
  122. # # print(sql)
  123. # cursor.execute(sql)
  124. # conn.commit()
  125. # cursor.close()
  126. # conn.close()
  127. # f.flush()
  128. # f.close()
  129. #print_metrics(history_model)
  130. def train():
  131. train_x,train_y,_ = getTokensLabels(isTrain=True,t="hand_label_person")
  132. test_x,test_y,test_context = getTokensLabels(isTrain=False,t="hand_label_person")
  133. with tf.Session() as sess:
  134. vocab,matrix = getVocabAndMatrix(getModel_w2v(),Embedding_size=128)
  135. model = getBiLSTMModel(input_shape=(2,10,128), vocab=vocab, embedding_weights=matrix, classes=4)
  136. callback = ModelCheckpoint(filepath="log/"+"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")
  137. model.fit(x=[train_x[0],train_x[1]],y=train_y,batch_size=128,epochs=600,callbacks=[callback],validation_data=[[test_x[0],test_x[1]],test_y])
  138. def predict():
  139. '''
  140. @summary: 预测数据
  141. '''
  142. def getTokensLabels():
  143. conn = psycopg2.connect(dbname="BidiPro",user="postgres",password="postgres",host="192.168.2.101")
  144. cursor = conn.cursor()
  145. #sql = '''
  146. #SELECT s.tokens,e.begin_index,e.end_index,e.doc_id,e.entity_id,e.sentence_index,e.entity_text,e.entity_type from entity_mention e,sentences s
  147. #WHERE s.doc_id=e.doc_id AND s.sentence_index=e.sentence_index AND e.entity_id not in (SELECT entity_id from entity_label) and entity_type in ('person') limit 10000
  148. #'''
  149. sql = '''
  150. SELECT s.tokens,e.begin_index,e.end_index,e.doc_id,e.entity_id,e.sentence_index,e.entity_text,e.entity_type from entity_mention e,sentences s
  151. WHERE s.doc_id=e.doc_id AND s.sentence_index=e.sentence_index AND e.doc_id in(SELECT doc_id from articles_validation) and entity_type in ('person')
  152. '''
  153. cursor.execute(sql)
  154. print(sql)
  155. data_x = []
  156. doc_id = []
  157. ent_id = []
  158. sen = []
  159. ent_text = []
  160. dianhua = []
  161. rows = cursor.fetchmany(1000)
  162. key_word = re.compile('电话[:|:]\d{7,12}|联系方式[:|:]\d{7,12}')
  163. phone = re.compile('1[3|4|5|7|8][0-9][-|——|—]?\d{4}[-|——|—]?\d{4}|\d{3,4}[-|——|—]\d{7,8}/\d{3,8}|\d{3,4}[-|——|—]\d{7,8}转\d{1,4}|\d{3,4}[-|——|—]\d{7,8}|[\(|\(]0\d{2,3}[\)|\)]\d{7,8}') # 联系电话
  164. while(rows):
  165. for row in rows:
  166. item_x = embedding(spanWindow(tokens=row[0],begin_index=row[1],end_index=row[2]))
  167. s = spanWindow(tokens=row[0],begin_index=row[1],end_index=row[2],size=15)
  168. s2 = ''.join(s[1])
  169. s2 = re.sub(',)', '-', s2)
  170. s2 = re.sub('\s','',s2)
  171. have_key = re.findall(key_word, s2)
  172. have_phone = re.findall(phone, s2)
  173. if have_phone:
  174. dianhua.append(have_phone)
  175. elif have_key:
  176. dianhua.append(have_phone)
  177. else:
  178. dianhua.append('')
  179. sen.append(s2)
  180. ent_id.append(row[4])
  181. ent_text.append(row[6])
  182. data_x.append(item_x)
  183. doc_id.append(row[3])
  184. rows = cursor.fetchmany(1000)
  185. cursor.close()
  186. conn.close()
  187. return np.transpose(np.array(data_x),(1,0,2,3)),doc_id,ent_id,sen,ent_text,dianhua
  188. test_x,doc_id,ent_id,sen,ent_text,dianhua = getTokensLabels()
  189. model = models.load_model("model_person.model",custom_objects={'precision':precision,'recall':recall,'f1_score':f1_score})
  190. predict_y = model.predict([test_x[0],test_x[1]])
  191. label = [np.argmax(y) for y in predict_y]
  192. data = {'doc_id':doc_id, 'ent_id':ent_id, 'sen':sen, 'entity_text':ent_text, 'dianhua':dianhua, 'label':label}
  193. df = pd.DataFrame(data)
  194. df.to_excel('data/person_phone.xls')
  195. conn = psycopg2.connect(dbname='BidiPro', user='postgres',password='postgres',host='192.168.2.101')
  196. cursor = conn.cursor()
  197. table = 'person_phone_predict'
  198. cursor.execute(" select to_regclass('"+table+"') is null ")
  199. notExists = cursor.fetchall()[0][0]
  200. if notExists:
  201. cursor.execute(" create table "+table+" (doc_id text,entity_id text,entity text,label int,predict text,phone text)")
  202. else:
  203. cursor.execute(" delete from "+table)
  204. for i in range(len(df['ent_id'])):
  205. pre_y = [str(a) for a in predict_y[i]]
  206. sql = " insert into "+table+"(doc_id,entity_id,entity,label,predict,phone) values('"+str(df['doc_id'][i])+"','"+str(df['ent_id'][i])+"','"+str(df['entity_text'][i])+"',"+str(int(df['label'][i]))+",'"+str(','.join(pre_y))+"','"+str(','.join(df['dianhua'][i]))+"')"
  207. #print(sql)
  208. cursor.execute(sql)
  209. conn.commit()
  210. print('提交完成')
  211. cursor.close()
  212. conn.close()
  213. if __name__ == '__main__':
  214. #get_data()
  215. #label_data()
  216. #post_data()
  217. training()
  218. predict()
  219. # train()