IterateModeling_LR.py 23 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302303304305306307308309310311312313314315316317318319320321322323324325326327328329330331332333334335336337338339340341342343344345346347348349350351352353354355356357358359360361362363364365366367368369370371372373374375376377378379380381382383384385386387388389390391392393394395396397398399400401402403404405406407408409410411412413414415416417418419420421422423424425426427428429430431432433434435436437438439440441442443444445446447448449450451452453454455456457458459460461462463464465466467468469470471472473474475476477478479480481482483484485486487488489490491492493494495496497498499500501502503504505506507508509510511512513514515516517518519520521522523524525526527528529530531532533534535536537538539540541542543544545546547548549550551552553554555556557558559560561562563564565566567568569570571
  1. # -*- coding: utf-8 -*-
  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. import glob
  10. sys.path.append(os.path.abspath("../.."))
  11. import psycopg2
  12. from keras import models
  13. from keras import layers
  14. from keras import optimizers,losses,metrics
  15. from keras.callbacks import ModelCheckpoint
  16. import codecs
  17. import copy
  18. from BiddingKG.dl.common.Utils import *
  19. import pandas as pd
  20. sourcetable = "label_guest_role"
  21. domain = sourcetable.split("_")[2]
  22. model_file = "model_"+domain+".model"
  23. input_shape = (2,10,128)
  24. output_shape = [6]
  25. def getTokensLabels(t,isTrain=True,predict=False):
  26. '''
  27. @param:
  28. t:标注数据所在表
  29. isTrain:是否训练
  30. predict:是否是验证
  31. @return:返回标注数据的处理后的输入和标签
  32. '''
  33. conn = psycopg2.connect(dbname="BiddingKM_test_10000",user="postgres",password="postgres",host="192.168.2.101")
  34. cursor = conn.cursor()
  35. if predict:
  36. sql = '''
  37. select A.tokens,B.begin_index,B.end_index,0,B.entity_id from sentences A,entity_mention_copy B where B.entity_type in ('org','company') and A.doc_id=B.doc_id and A.sentence_index=B.sentence_index
  38. and B.doc_id in (select doc_id from articles_validation ) order by B.doc_id
  39. '''
  40. else:
  41. select_sql = " select A.tokens,B.begin_index,B.end_index,C.label,C.entity_id "
  42. '''
  43. if isTrain:
  44. train_sql = " and C.id not in(select variable_id from dd_graph_variables_holdout) "
  45. else:
  46. train_sql = " and C.id in(select variable_id from dd_graph_variables_holdout)"
  47. '''
  48. '''
  49. if isTrain:
  50. train_sql = " and A.doc_id not in(select id from articles_processed order by id limit 1000) "
  51. else:
  52. train_sql = " and A.doc_id in(select id from articles_processed order by id limit 1000)"
  53. '''
  54. if isTrain:
  55. train_sql = " and C.entity_id not in(select entity_id from is_wintenderer_label_inference where id in(select variable_id from dd_graph_variables_holdout))"
  56. else:
  57. #train_sql = " and C.entity_id in(select entity_id from is_wintenderer_label_inference where id in(select variable_id from dd_graph_variables_holdout))"
  58. train_sql = " and exists(select 1 from test_predict_money h,entity_mention g where h.entity_id=g.entity_id and A.doc_id=g.doc_id) order by B.doc_id limit 2000 "
  59. sql = select_sql+" from sentences A,entity_mention_copy B,"+t+" C where B.entity_type in ('org','company') and A.doc_id=B.doc_id and A.sentence_index=B.sentence_index and B.entity_id=C.entity_id "+train_sql
  60. print(sql)
  61. cursor.execute(sql)
  62. data_x = []
  63. data_y = []
  64. data_context = []
  65. rows = cursor.fetchmany(1000)
  66. allLimit = 330000
  67. all = 0
  68. while(rows):
  69. for row in rows:
  70. if all>=allLimit:
  71. break
  72. item_x = embedding(spanWindow(tokens=row[0],begin_index=row[1],end_index=row[2],size=input_shape[1]),shape=input_shape)
  73. item_y = np.zeros(output_shape)
  74. item_y[row[3]] = 1
  75. all += 1
  76. if not isTrain:
  77. item_context = []
  78. item_context.append(row[4])
  79. data_context.append(item_context)
  80. data_x.append(item_x)
  81. data_y.append(item_y)
  82. rows = cursor.fetchmany(1000)
  83. return np.transpose(np.array(data_x),(1,0,2,3)),np.array(data_y),data_context
  84. def getBiRNNModel():
  85. '''
  86. @summary:获取模型
  87. '''
  88. L_input = layers.Input(shape=input_shape[1:],dtype="float32")
  89. #C_input = layers.Input(shape=(10,128),dtype="float32")
  90. R_input = layers.Input(shape=input_shape[1:],dtype="float32")
  91. #lstm_0 = layers.Bidirectional(layers.LSTM(16,return_sequences=True))(ThreeBilstm(0)(input))
  92. lstm_0 = layers.Bidirectional(layers.LSTM(16,return_sequences=True))(L_input)
  93. avg_0 = layers.GlobalAveragePooling1D()(lstm_0)
  94. #lstm_1 = layers.Bidirectional(layers.LSTM(16,return_sequences=True))(C_input)
  95. #avg_1 = layers.GlobalAveragePooling1D()(lstm_1)
  96. lstm_2 = layers.Bidirectional(layers.LSTM(16,return_sequences=True))(R_input)
  97. avg_2 = layers.GlobalAveragePooling1D()(lstm_2)
  98. #concat = layers.merge([avg_0,avg_1,avg_2],mode="concat")
  99. concat = layers.merge([avg_0,avg_2],mode="concat")
  100. output = layers.Dense(output_shape[0],activation="softmax")(concat)
  101. model = models.Model(inputs=[L_input,R_input],outputs=output)
  102. model.compile(optimizer=optimizers.Adam(lr=0.001),loss=losses.binary_crossentropy,metrics=[precision,recall,f1_score])
  103. return model
  104. def loadTrainData(percent=0.9):
  105. files = ["id_token_text_begin_end_label.pk","id_token_text_begin_end_label.pk1"]
  106. data_x = []
  107. data_y = []
  108. #data_id = []
  109. test_x = []
  110. test_y = []
  111. test_id = []
  112. for file in files:
  113. data = load(file)
  114. for row in data:
  115. item_x = embedding(spanWindow(tokens=row[1],begin_index=row[3],end_index=row[4],size=input_shape[1]),shape=input_shape)
  116. item_y = np.zeros(output_shape)
  117. label = int(row[5])
  118. if label not in [0,1,2,3,4,5]:
  119. continue
  120. item_y[label] = 1
  121. if np.random.random()<percent:
  122. data_x.append(item_x)
  123. data_y.append(item_y)
  124. #data_id.append(row[0])
  125. else:
  126. test_x.append(item_x)
  127. test_y.append(item_y)
  128. test_id.append(row[0])
  129. return np.transpose(np.array(data_x),(1,0,2,3)),np.array(data_y),np.transpose(np.array(test_x),(1,0,2,3)),np.array(test_y),None,test_id
  130. def training():
  131. '''
  132. @summary:训练模型
  133. '''
  134. model = getBiRNNModel()
  135. model.summary()
  136. #train_x,train_y,_ = getTokensLabels(isTrain=True,t="hand_label_role")
  137. train_x,train_y,test_x,test_y,_,test_id = loadTrainData()
  138. save([test_x,test_y,test_id],"val_data.pk")
  139. checkpoint = ModelCheckpoint(
  140. "../../dl_dev/role/log/ep{epoch:03d}-loss{loss:.3f}-val_loss{val_loss:.3f}-f1{val_f1_score:.3f}.h5", monitor="val_loss", verbose=1, save_best_only=True, mode='min')
  141. print(np.shape(train_x))
  142. 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=120,batch_size=512,shuffle=True,callbacks=[checkpoint])
  143. predict_y = model.predict([test_x[0],test_x[1]])
  144. model.save(model_file)
  145. #print_metrics(history_model)
  146. def val():
  147. files = []
  148. for file in glob.glob("C:\\Users\\User\\Desktop\\20190416要素\\*.html"):
  149. filename = file.split("\\")[-1]
  150. files.append(filename)
  151. conn = psycopg2.connect(dbname="article_label",user="postgres",password="postgres",host="192.168.2.101")
  152. cursor = conn.cursor()
  153. sql = '''
  154. select A.entity_id,A.entity_text,A.begin_index,A.end_index,A.label,A.values,B.tokens,A.doc_id
  155. from entity_mention A,sentences B
  156. where A.doc_id=B.doc_id and A.sentence_index=B.sentence_index
  157. and A.entity_type in ('org','company')
  158. and A.label!='None'
  159. and not exists(select 1 from turn_label where entity_id=A.entity_id)
  160. order by A.label
  161. '''
  162. cursor.execute(sql)
  163. rows = cursor.fetchall()
  164. list_entity_id = []
  165. list_before = []
  166. list_after = []
  167. list_text = []
  168. list_label = []
  169. list_prob = []
  170. repeat = set()
  171. data_x = []
  172. cnn_x = []
  173. for row in rows:
  174. entity_id = row[0]
  175. entity_text = row[1]
  176. begin_index = row[2]
  177. end_index = row[3]
  178. label = int(row[4])
  179. values = row[5][1:-1].split(",")
  180. tokens = row[6]
  181. doc_id = row[7]
  182. if doc_id not in files:
  183. continue
  184. if float(values[label])<0.5:
  185. continue
  186. beforeafter = spanWindow(tokens, begin_index, end_index, 10,center_include=True,text=entity_text)
  187. if ("".join(beforeafter[0]),"".join(beforeafter[1]),"".join(beforeafter[2])) in repeat:
  188. continue
  189. repeat.add(("".join(beforeafter[0]),"".join(beforeafter[1]),"".join(beforeafter[2])))
  190. item_x = embedding(spanWindow(tokens=tokens,begin_index=begin_index,end_index=end_index,size=input_shape[1]),shape=input_shape)
  191. data_x.append(item_x)
  192. cnn_x.append(encodeInput(spanWindow(tokens=tokens,begin_index=begin_index,end_index=end_index,size=10,center_include=True,word_flag=True,text=entity_text), word_len=50, word_flag=True))
  193. list_entity_id.append(entity_id)
  194. list_before.append("".join(beforeafter[0]))
  195. list_after.append("".join(beforeafter[2]))
  196. list_text.append("".join(beforeafter[1]))
  197. list_label.append(label)
  198. list_prob.append(values[label])
  199. model = models.load_model("../../dl_dev/role/log/new_biLSTM-ep012-loss0.028-val_loss0.040-f10.954.h5", custom_objects={"precision":precision, "recall":recall, "f1_score":f1_score})
  200. data_x = np.transpose(np.array(data_x),(1,0,2,3))
  201. predict_value = model.predict([data_x[0],data_x[1]])
  202. predict_y = np.argmax(predict_value,1)
  203. list_newprob = []
  204. for label,value in zip(predict_y,predict_value):
  205. list_newprob.append(value[label])
  206. print("len",len(list_entity_id))
  207. model_cnn = models.load_model("../../dl_dev/role/log/ep071-loss0.107-val_loss0.122-f10.956.h5", custom_objects={"precision":precision, "recall":recall, "f1_score":f1_score})
  208. cnn_x = np.transpose(np.array(cnn_x),(1,0,2))
  209. predict_value = model_cnn.predict([cnn_x[0],cnn_x[1],cnn_x[2]])
  210. predict_y_cnn = np.argmax(predict_value,1)
  211. list_newprob_cnn = []
  212. for label,value in zip(predict_y_cnn,predict_value):
  213. list_newprob_cnn.append(value[label])
  214. print("len",len(list_entity_id))
  215. data = []
  216. for id,before,text,after,label_bi,prob,label_bi1,newprob,label_cnn,newprob_cnn in zip(list_entity_id,list_before,list_text,list_after,list_label,list_prob,predict_y,list_newprob,predict_y_cnn,list_newprob_cnn):
  217. if label_bi1!=label_cnn:
  218. data.append([id,before,text,after,label_bi,prob,label_bi1,newprob,label_cnn,newprob_cnn])
  219. data.sort(key=lambda x:x[6])
  220. list_entity_id = []
  221. list_before = []
  222. list_after = []
  223. list_text = []
  224. list_label = []
  225. list_prob = []
  226. list_newlabel = []
  227. list_newprob = []
  228. list_newlabel_cnn = []
  229. list_newprob_cnn = []
  230. for item in data:
  231. list_entity_id.append(item[0])
  232. list_before.append(item[1])
  233. list_text.append(item[2])
  234. list_after.append(item[3])
  235. list_label.append(item[4])
  236. list_prob.append(item[5])
  237. list_newlabel.append(item[6])
  238. list_newprob.append(item[7])
  239. list_newlabel_cnn.append(item[8])
  240. list_newprob_cnn.append(item[9])
  241. parts = 1
  242. parts_num = len(list_entity_id)//parts
  243. for i in range(parts-1):
  244. data = {"entity_id":list_entity_id[i*parts_num:(i+1)*parts_num],"list_before":list_before[i*parts_num:(i+1)*parts_num],"list_after":list_after[i*parts_num:(i+1)*parts_num],"list_text":list_text[i*parts_num:(i+1)*parts_num],"list_label":list_label[i*parts_num:(i+1)*parts_num],"list_prob":list_prob[i*parts_num:(i+1)*parts_num]}
  245. df = pd.DataFrame(data)
  246. df.to_excel("未标注错误_"+str(i)+".xls",columns=["entity_id","list_before","list_text","list_after","list_label","list_prob"])
  247. i = parts - 1
  248. data = {"entity_id":list_entity_id[i*parts_num:],"list_before":list_before[i*parts_num:],"list_after":list_after[i*parts_num:],"list_text":list_text[i*parts_num:],"list_label":list_label[i*parts_num:],"list_prob":list_prob[i*parts_num:],"list_newlabel":list_newlabel[i*parts_num:],"list_newprob":list_newprob[i*parts_num:],"list_newlabel_cnn":list_newlabel_cnn[i*parts_num:],"list_newprob_cnn":list_newprob_cnn[i*parts_num:]}
  249. df = pd.DataFrame(data)
  250. df.to_excel("测试数据_role-cnnw-biw"+str(i)+".xls",columns=["entity_id","list_before","list_text","list_after","list_label","list_prob","list_newlabel","list_newprob","list_newlabel_cnn","list_newprob_cnn"])
  251. def validation():
  252. conn1 = psycopg2.connect(dbname="article_label",user="postgres",password="postgres",host="192.168.2.101")
  253. conn2 = psycopg2.connect(dbname="BiddingKM_test_10000",user="postgres",password="postgres",host="192.168.2.101")
  254. cursor1 = conn1.cursor()
  255. cursor2 = conn2.cursor()
  256. model = getBiRNNModel()
  257. model.load_weights("log/ep010-loss0.033-val_loss0.043-f10.950.h5")
  258. [test_x,test_y,test_id] = load("val_data.pk")
  259. predict_y = model.predict([test_x[0],test_x[1]])
  260. list_id = []
  261. list_before = []
  262. list_text = []
  263. list_after = []
  264. list_same = []
  265. list_predict = []
  266. list_label = []
  267. data = []
  268. for id,predict,label in zip(test_id,np.argmax(predict_y,1),np.argmax(test_y,1)):
  269. if predict==label:
  270. same = 0
  271. text = ""
  272. beforeafter = [[],[]]
  273. else:
  274. same = 1
  275. if re.search("比地",id) is not None:
  276. sql = " select A.tokens,B.entity_text,B.begin_index,B.end_index from sentences A,entity_mention B where A.doc_id=B.doc_id and A.sentence_index=B.sentence_index and B.entity_id='"+id+"' "
  277. cursor1.execute(sql)
  278. rows = cursor1.fetchall()
  279. else:
  280. sql = " select A.tokens,B.entity_text,B.begin_index,B.end_index from sentences A,entity_mention_copy B where A.doc_id=B.doc_id and A.sentence_index=B.sentence_index and B.entity_id='"+id+"' "
  281. cursor2.execute(sql)
  282. rows = cursor2.fetchall()
  283. retu = rows[0]
  284. text = retu[1]
  285. beforeafter = spanWindow(retu[0], retu[2], retu[3], 10)
  286. data.append([id,same,"".join(beforeafter[0]),text,"".join(beforeafter[1]),label,predict])
  287. data.sort(key=lambda x:x[1])
  288. for item in data:
  289. list_id.append(item[0])
  290. list_same.append(item[1])
  291. list_before.append(item[2])
  292. list_text.append(item[3])
  293. list_after.append(item[4])
  294. list_label.append(item[5])
  295. list_predict.append(item[6])
  296. df = pd.DataFrame({"list_id":list_id,"list_same":list_same,"list_before":list_before,"list_text":list_text,"list_after":list_after,"list_label":list_label,"list_predict":list_predict})
  297. columns = ["list_id","list_same","list_before","list_text","list_after","list_label","list_predict"]
  298. df.to_excel("result.xls",index=False,columns=columns)
  299. conn1.close()
  300. conn2.close()
  301. def trainingIteration_category(iterate=2,label_table=sourcetable):
  302. '''
  303. @summary: 迭代训练模型,修改标签,适用于当数据准确率不高的条件
  304. @param:
  305. iterate:迭代次数
  306. label_table:标签数据所在表
  307. '''
  308. def getDatasets():
  309. conn = psycopg2.connect(dbname="BiddingKM_test_10000",user="postgres",password="postgres",host="192.168.2.101")
  310. cursor = conn.cursor()
  311. select_sql = " select A.tokens,B.begin_index,B.end_index,C.label,C.entity_id "
  312. sql = select_sql+" from sentences A,entity_mention B,"+label_table+" C where A.doc_id=B.doc_id and A.sentence_index=B.sentence_index and B.entity_id=C.entity_id order by A.doc_id "
  313. cursor.execute(sql)
  314. print(sql)
  315. data_x = []
  316. data_y = []
  317. id_set = []
  318. rows = cursor.fetchmany(1000)
  319. allLimit = 320000
  320. all = 0
  321. while(rows):
  322. for row in rows:
  323. if all>=allLimit:
  324. break
  325. item_x = embedding(spanWindow(tokens=row[0],begin_index=row[1],end_index=row[2]))
  326. item_y = np.zeros(output_shape)
  327. item_y[row[3]] = 1
  328. all += 1
  329. data_x.append(item_x)
  330. data_y.append(item_y)
  331. id_set.append(row[4])
  332. rows = cursor.fetchmany(1000)
  333. return np.transpose(np.array(data_x),(1,0,2,3)),np.array(data_y),id_set
  334. train_x,train_y,id_set = getDatasets()
  335. alllength = len(train_x[0])
  336. parts = 6
  337. num_parts = alllength//parts
  338. copy_y = copy.copy(train_y)
  339. for ite in range(iterate):
  340. for j in range(parts-1):
  341. print("iterate:",str(ite)+"/"+str(iterate-1),str(j)+"/"+str(parts-1))
  342. model = getBiRNNModel()
  343. model.summary()
  344. test_begin = j*num_parts
  345. test_end = (j+1)*num_parts
  346. checkpoint = ModelCheckpoint(model_file+".hdf5",monitor="val_loss",verbose=1,save_best_only=True,mode='min')
  347. history_model = model.fit(x=[np.concatenate((train_x[0][0:test_begin],train_x[0][test_end:])),np.concatenate((train_x[1][0:test_begin],train_x[1][test_end:]))],y=np.concatenate((copy_y[0:test_begin],copy_y[test_end:])),validation_data=([train_x[0][test_begin:test_end],train_x[1][test_begin:test_end]],copy_y[test_begin:test_end]),epochs=30,batch_size=300,shuffle=True,callbacks=[checkpoint])
  348. model.load_weights(model_file+".hdf5")
  349. predict_y = model.predict([train_x[0][test_begin:test_end],train_x[1][test_begin:test_end]])
  350. for i in range(len(predict_y)):
  351. if np.max(predict_y[i])>=0.8:
  352. max_index = np.argmax(predict_y[i])
  353. for h in range(len(predict_y[i])):
  354. if h==max_index:
  355. copy_y[i+test_begin][h] = 1
  356. else:
  357. copy_y[i+test_begin][h] = 0
  358. print("iterate:",str(ite)+"/"+str(iterate-1),str(j)+"/"+str(parts-1))
  359. model = getBiRNNModel()
  360. model.summary()
  361. test_begin = j*num_parts
  362. checkpoint = ModelCheckpoint(model_file+".hdf5",monitor="val_loss",verbose=1,save_best_only=True,mode="min")
  363. history_model = model.fit(x=[train_x[0][0:test_begin],train_x[1][0:test_begin]],y=copy_y[0:test_begin],validation_data=([train_x[0][test_begin:],train_x[1][test_begin:]],copy_y[test_begin:]),epochs=30,batch_size=300,shuffle=True,callbacks=[checkpoint])
  364. model.load_weights(model_file+".hdf5")
  365. predict_y = model.predict([train_x[0][test_begin:],train_x[1][test_begin:]])
  366. for i in range(len(predict_y)):
  367. if np.max(predict_y[i])>=0.8:
  368. max_index = np.argmax(predict_y[i])
  369. for h in range(len(predict_y[i])):
  370. if h==max_index:
  371. copy_y[i+test_begin][h] = 1
  372. else:
  373. copy_y[i+test_begin][h] = 0
  374. with codecs.open("final_label_"+domain+".txt","w",encoding="utf8") as f:
  375. for i in range(len(id_set)):
  376. f.write(id_set[i])
  377. f.write("\t")
  378. f.write(str(np.argmax(copy_y[i])))
  379. f.write("\n")
  380. f.flush()
  381. f.close()
  382. def predict():
  383. '''
  384. @summary: 预测测试数据
  385. '''
  386. test_x,_,ids = getTokensLabels("final_label_role", isTrain=False,predict=True)
  387. model = models.load_model(model_file,custom_objects={'precision':precision,'recall':recall,'f1_score':f1_score})
  388. predict_y = model.predict([test_x[0],test_x[1]])
  389. with codecs.open("test_predict_"+domain+".txt","w",encoding="utf8") as f:
  390. for i in range(len(predict_y)):
  391. f.write(ids[i][0])
  392. f.write("\t")
  393. f.write(str(np.argmax(predict_y[i])))
  394. f.write("\t")
  395. value = ""
  396. for item in predict_y[i]:
  397. value += str(item)+","
  398. f.write(value[:-1])
  399. f.write("\n")
  400. f.flush()
  401. f.close()
  402. def importIterateLabel():
  403. '''
  404. @summary:导入迭代之后的标签值
  405. '''
  406. file = "final_label_"+domain+".txt"
  407. conn = psycopg2.connect(dbname="BiddingKM_test_10000",user="postgres",password="postgres",host="192.168.2.101")
  408. cursor = conn.cursor()
  409. tablename = file.split(".")[0]
  410. # 创建表
  411. cursor.execute(" SELECT to_regclass('"+tablename+"') is null ")
  412. flag = cursor.fetchall()[0][0]
  413. if flag:
  414. cursor.execute(" create table "+tablename+"(entity_id text,label int)")
  415. else:
  416. cursor.execute(" delete from "+tablename)
  417. with codecs.open(file,"r",encoding="utf8") as f:
  418. while(True):
  419. line = f.readline()
  420. if not line:
  421. break
  422. line_split = line.split("\t")
  423. entity_id=line_split[0]
  424. label = line_split[1]
  425. sql = " insert into "+tablename+"(entity_id,label) values('"+str(entity_id)+"',"+str(label)+")"
  426. cursor.execute(sql)
  427. f.close()
  428. conn.commit()
  429. conn.close()
  430. def importtestPredict():
  431. '''
  432. @summary:导入测试数据的预测值
  433. '''
  434. file = "test_predict_"+domain+".txt"
  435. conn = psycopg2.connect(dbname="BiddingKG",user="postgres",password="postgres",host="192.168.2.101")
  436. cursor = conn.cursor()
  437. tablename = file.split(".")[0]
  438. # 创建表
  439. cursor.execute(" SELECT to_regclass('"+tablename+"') is null ")
  440. flag = cursor.fetchall()[0][0]
  441. if flag:
  442. cursor.execute(" create table "+tablename+"(entity_id text,label int,value text)")
  443. else:
  444. cursor.execute(" delete from "+tablename)
  445. with codecs.open(file,"r",encoding="utf8") as f:
  446. while(True):
  447. line = f.readline()
  448. if not line:
  449. break
  450. line_split = line.split("\t")
  451. entity_id=line_split[0]
  452. predict = line_split[1]
  453. value = line_split[2]
  454. sql = " insert into "+tablename+"(entity_id,label,value) values('"+str(entity_id)+"',"+str(predict)+",'"+str(value)+"')"
  455. cursor.execute(sql)
  456. f.close()
  457. conn.commit()
  458. conn.close()
  459. def autoIterate():
  460. #trainingIteration_binary()
  461. trainingIteration_category()
  462. importIterateLabel()
  463. training()
  464. predict()
  465. def test1(entity_id):
  466. conn = psycopg2.connect(dbname="article_label",user="postgres",password="postgres",host="192.168.2.101")
  467. cursor = conn.cursor()
  468. if predict:
  469. sql = "select A.tokens,B.begin_index,B.end_index,0,B.entity_id from sentences A,entity_mention B where B.entity_type in ('org','company') and A.doc_id=B.doc_id and A.sentence_index=B.sentence_index and B.entity_id='"+entity_id+"'"
  470. print(sql)
  471. cursor.execute(sql)
  472. data_x = []
  473. data_y = []
  474. rows = cursor.fetchmany(1000)
  475. while(rows):
  476. for row in rows:
  477. item_x = encodeInput(spanWindow(tokens=row[0],begin_index=row[1],end_index=row[2],size=10,center_include=True,word_flag=True), word_len=50, word_flag=True)
  478. item_y = np.zeros(output_shape)
  479. item_y[row[3]] = 1
  480. data_x.append(item_x)
  481. data_y.append(item_y)
  482. rows = cursor.fetchmany(1000)
  483. model = models.load_model("../../dl_dev/role/log/ep017-loss0.088-val_loss0.125-f10.955.h5", custom_objects={'precision':precision, 'recall':recall, 'f1_score':f1_score})
  484. test_x = np.transpose(np.array(data_x),(1,0,2))
  485. predict_y = model.predict([test_x[0],test_x[1],test_x[2]])
  486. print(predict_y)
  487. if __name__=="__main__":
  488. #training()
  489. val()
  490. #validation()
  491. #test()
  492. #trainingIteration_category()
  493. #importIterateLabel()
  494. #predict()
  495. #importtestPredict()
  496. #autoIterate()
  497. #test1("比地_101_61333318.html_0_116_122")