generateData.py 30 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302303304305306307308309310311312313314315316317318319320321322323324325326327328329330331332333334335336337338339340341342343344345346347348349350351352353354355356357358359360361362363364365366367368369370371372373374375376377378379380381382383384385386387388389390391392393394395396397398399400401402403404405406407408409410411412413414415416417418419420421422423424425426427428429430431432433434435436437438439440441442443444445446447448449450451452453454455456457458459460461462463464465466467468469470471472473474475476477478479480481482483484485486487488489490491492493494495496497498499500501502503504505506507508509510511512513514515516517518519520521522523524525526527528529530531532533534535536537538539540541542543544545546547548549550551552553554555556557558559560561562563564565566567568569570571572573574575576577578579580581582583584585586587588589590591592593594595596597598599600601602603604605606607608609610611612613614615616617618619620621622623624625626627628629630631632633634635636637638639640641642643644645646647648649650651652653654655656657658659660661662663664665666667668669670671672673674675676677678679680681682683684685686687688689690691692693694695696697698699700701702703704705706707708709710711712713714715716717718719720721
  1. '''
  2. Created on 2019年3月25日
  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 glob
  11. import re
  12. import copy
  13. from bs4 import BeautifulSoup
  14. import codecs
  15. import pandas as pd
  16. from BiddingKG.dl.interface.predictor import *
  17. from BiddingKG.dl.form.feature import *
  18. import psycopg2
  19. from BiddingKG.dl.common.Utils import *
  20. # formPredictor = FormPredictor()
  21. def tableToText(soup,data,file,data_set_is,data_set_no):
  22. '''
  23. @param:
  24. soup:网页html的soup
  25. @return:处理完表格信息的网页text
  26. '''
  27. def getTrs(tbody):
  28. #获取所有的tr
  29. trs = []
  30. objs = tbody.find_all(recursive=False)
  31. for obj in objs:
  32. if obj.name=="tr":
  33. trs.append(obj)
  34. if obj.name=="tbody":
  35. for tr in obj.find_all("tr",recursive=False):
  36. trs.append(tr)
  37. return trs
  38. def fixSpan(tbody):
  39. # 处理colspan, rowspan信息补全问题
  40. #trs = tbody.findChildren('tr', recursive=False)
  41. trs = getTrs(tbody)
  42. ths_len = 0
  43. ths = list()
  44. trs_set = set()
  45. #修改为先进行列补全再进行行补全,否则可能会出现表格解析混乱
  46. # 遍历每一个tr
  47. for indtr, tr in enumerate(trs):
  48. ths_tmp = tr.findChildren('th', recursive=False)
  49. #不补全含有表格的tr
  50. if len(tr.findChildren('table'))>0:
  51. continue
  52. if len(ths_tmp) > 0:
  53. ths_len = ths_len + len(ths_tmp)
  54. for th in ths_tmp:
  55. ths.append(th)
  56. trs_set.add(tr)
  57. # 遍历每行中的element
  58. tds = tr.findChildren(recursive=False)
  59. for indtd, td in enumerate(tds):
  60. # 若有colspan 则补全同一行下一个位置
  61. if 'colspan' in td.attrs:
  62. if str(re.sub("[^0-9]","",str(td['colspan'])))!="":
  63. col = int(re.sub("[^0-9]","",str(td['colspan'])))
  64. td['colspan'] = 1
  65. for i in range(1, col, 1):
  66. td.insert_after(copy.copy(td))
  67. for indtr, tr in enumerate(trs):
  68. ths_tmp = tr.findChildren('th', recursive=False)
  69. #不补全含有表格的tr
  70. if len(tr.findChildren('table'))>0:
  71. continue
  72. if len(ths_tmp) > 0:
  73. ths_len = ths_len + len(ths_tmp)
  74. for th in ths_tmp:
  75. ths.append(th)
  76. trs_set.add(tr)
  77. # 遍历每行中的element
  78. tds = tr.findChildren(recursive=False)
  79. for indtd, td in enumerate(tds):
  80. # 若有rowspan 则补全下一行同样位置
  81. if 'rowspan' in td.attrs:
  82. if str(re.sub("[^0-9]","",str(td['rowspan'])))!="":
  83. row = int(re.sub("[^0-9]","",str(td['rowspan'])))
  84. td['rowspan'] = 1
  85. for i in range(1, row, 1):
  86. # 获取下一行的所有td, 在对应的位置插入
  87. if indtr+i<len(trs):
  88. tds1 = trs[indtr + i].findChildren(['td','th'], recursive=False)
  89. if len(tds1) >= (indtd) and len(tds1)>0:
  90. if indtd > 0:
  91. tds1[indtd - 1].insert_after(copy.copy(td))
  92. else:
  93. tds1[0].insert_before(copy.copy(td))
  94. def getTable(tbody):
  95. #trs = tbody.findChildren('tr', recursive=False)
  96. trs = getTrs(tbody)
  97. inner_table = []
  98. for tr in trs:
  99. tr_line = []
  100. tds = tr.findChildren(['td','th'], recursive=False)
  101. for td in tds:
  102. tr_line.append([re.sub('\s*','',td.get_text()),0])
  103. inner_table.append(tr_line)
  104. return inner_table
  105. #处理表格不对齐的问题
  106. def fixTable(inner_table):
  107. maxWidth = 0
  108. for item in inner_table:
  109. if len(item)>maxWidth:
  110. maxWidth = len(item)
  111. for i in range(len(inner_table)):
  112. if len(inner_table[i])<maxWidth:
  113. for j in range(maxWidth-len(inner_table[i])):
  114. inner_table[i].append(["",0])
  115. return inner_table
  116. def removePadding(inner_table,pad_row = "@@",pad_col = "##"):
  117. height = len(inner_table)
  118. width = len(inner_table[0])
  119. for i in range(height):
  120. point = ""
  121. for j in range(width):
  122. if inner_table[i][j][0]==point and point!="":
  123. inner_table[i][j][0] = pad_row
  124. else:
  125. if inner_table[i][j][0] not in [pad_row,pad_col]:
  126. point = inner_table[i][j][0]
  127. for j in range(width):
  128. point = ""
  129. for i in range(height):
  130. if inner_table[i][j][0]==point and point!="":
  131. inner_table[i][j][0] = pad_col
  132. else:
  133. if inner_table[i][j][0] not in [pad_row,pad_col]:
  134. point = inner_table[i][j][0]
  135. def addPadding(inner_table,pad_row = "@@",pad_col = "##"):
  136. height = len(inner_table)
  137. width = len(inner_table[0])
  138. for i in range(height):
  139. for j in range(width):
  140. if inner_table[i][j][0]==pad_row:
  141. inner_table[i][j][0] = inner_table[i][j-1][0]
  142. inner_table[i][j][1] = inner_table[i][j-1][1]
  143. if inner_table[i][j][0]==pad_col:
  144. inner_table[i][j][0] = inner_table[i-1][j][0]
  145. inner_table[i][j][1] = inner_table[i-1][j][1]
  146. #设置表头
  147. def setHead(inner_table,prob_min=0.64):
  148. pad_row = "@@"
  149. pad_col = "##"
  150. removePadding(inner_table, pad_row, pad_col)
  151. pad_pattern = re.compile(pad_row+"|"+pad_col)
  152. height = len(inner_table)
  153. width = len(inner_table[0])
  154. head_list = []
  155. head_list.append(0)
  156. #行表头
  157. is_head_last = False
  158. for i in range(height):
  159. is_head = False
  160. is_long_value = False
  161. #判断是否是全padding值
  162. is_same_value = True
  163. same_value = inner_table[i][0][0]
  164. for j in range(width):
  165. if inner_table[i][j][0]!=same_value and inner_table[i][j][0]!=pad_row:
  166. is_same_value = False
  167. break
  168. #predict is head or not with model
  169. temp_item = ""
  170. for j in range(width):
  171. temp_item += inner_table[i][j][0]+"|"
  172. temp_item = re.sub(pad_pattern,"",temp_item)
  173. form_prob = formPredictor.predict(encoding(temp_item,expand=True))
  174. if form_prob is not None:
  175. if form_prob[0][1]>prob_min:
  176. is_head = True
  177. else:
  178. is_head = False
  179. #print(temp_item,form_prob)
  180. if len(inner_table[i][0][0])>40:
  181. is_long_value = True
  182. if is_head or is_long_value or is_same_value:
  183. #不把连续表头分开
  184. if not is_head_last:
  185. head_list.append(i)
  186. if is_long_value or is_same_value:
  187. head_list.append(i+1)
  188. if is_head:
  189. for j in range(width):
  190. if inner_table[i][j][0] not in data_set_is and inner_table[i][j][0] not in data_set_no:
  191. data.append([file,inner_table[i][j][0],1])
  192. data_set_is.add(inner_table[i][j][0])
  193. inner_table[i][j][1] = 1
  194. is_head_last = is_head
  195. head_list.append(height)
  196. #列表头
  197. for i in range(len(head_list)-1):
  198. head_begin = head_list[i]
  199. head_end = head_list[i+1]
  200. #最后一列不设置为列表头
  201. for i in range(width-1):
  202. is_head = False
  203. #predict is head or not with model
  204. temp_item = ""
  205. for j in range(head_begin,head_end):
  206. temp_item += inner_table[j][i][0]+"|"
  207. temp_item = re.sub(pad_pattern,"",temp_item)
  208. form_prob = formPredictor.predict(encoding(temp_item,expand=True))
  209. if form_prob is not None:
  210. if form_prob[0][1]>prob_min:
  211. is_head = True
  212. else:
  213. is_head = False
  214. if is_head:
  215. for j in range(head_begin,head_end):
  216. if inner_table[j][i][0] not in data_set_is and inner_table[j][i][0] not in data_set_no:
  217. data.append([file,inner_table[j][i][0],1])
  218. data_set_is.add(inner_table[j][i][0])
  219. inner_table[j][i][1] = 2
  220. for line in inner_table:
  221. for item in line:
  222. if item[0] not in data_set_is and item[0] not in data_set_no:
  223. data.append([file,item[0],0])
  224. data_set_no.add(item[0])
  225. addPadding(inner_table, pad_row, pad_col)
  226. return inner_table,head_list
  227. #设置表头
  228. def setHead_withRule(inner_table,pattern,pat_value,count):
  229. height = len(inner_table)
  230. width = len(inner_table[0])
  231. head_list = []
  232. head_list.append(0)
  233. #行表头
  234. is_head_last = False
  235. for i in range(height):
  236. set_match = set()
  237. is_head = False
  238. is_long_value = False
  239. is_same_value = True
  240. same_value = inner_table[i][0][0]
  241. for j in range(width):
  242. if inner_table[i][j][0]!=same_value:
  243. is_same_value = False
  244. break
  245. for j in range(width):
  246. if re.search(pat_value,inner_table[i][j][0]) is not None:
  247. is_head = False
  248. break
  249. str_find = re.findall(pattern,inner_table[i][j][0])
  250. if len(str_find)>0:
  251. set_match.add(inner_table[i][j][0])
  252. if len(set_match)>=count:
  253. is_head = True
  254. if len(inner_table[i][0][0])>40:
  255. is_long_value = True
  256. if is_head or is_long_value or is_same_value:
  257. if not is_head_last:
  258. head_list.append(i)
  259. if is_head:
  260. for j in range(width):
  261. inner_table[i][j][1] = 1
  262. is_head_last = is_head
  263. head_list.append(height)
  264. #列表头
  265. for i in range(len(head_list)-1):
  266. head_begin = head_list[i]
  267. head_end = head_list[i+1]
  268. #最后一列不设置为列表头
  269. for i in range(width-1):
  270. set_match = set()
  271. is_head = False
  272. for j in range(head_begin,head_end):
  273. if re.search(pat_value,inner_table[j][i][0]) is not None:
  274. is_head = False
  275. break
  276. str_find = re.findall(pattern,inner_table[j][i][0])
  277. if len(str_find)>0:
  278. set_match.add(inner_table[j][i][0])
  279. if len(set_match)>=count:
  280. is_head = True
  281. if is_head:
  282. for j in range(head_begin,head_end):
  283. inner_table[j][i][1] = 2
  284. return inner_table,head_list
  285. #取得表格的处理方向
  286. def getDirect(inner_table,begin,end):
  287. column_head = set()
  288. row_head = set()
  289. widths = len(inner_table[0])
  290. for height in range(begin,end):
  291. for width in range(widths):
  292. if inner_table[height][width][1] ==1:
  293. row_head.add(height)
  294. if inner_table[height][width][1] ==2:
  295. column_head.add(width)
  296. company_pattern = re.compile("公司")
  297. if 0 in column_head and begin not in row_head:
  298. return "column"
  299. if 0 in column_head and begin in row_head:
  300. for height in range(begin,end):
  301. count = 0
  302. count_flag = True
  303. for width_index in range(width):
  304. if inner_table[height][width_index][1]==0:
  305. if re.search(company_pattern,inner_table[height][width_index][0]) is not None:
  306. count += 1
  307. else:
  308. count_flag = False
  309. if count_flag and count>=2:
  310. return "column"
  311. return "row"
  312. #根据表格处理方向生成句子,
  313. def getTableText(inner_table,head_list):
  314. rankPattern = "(排名|排序|名次|评标结果|评审结果)"
  315. entityPattern = "(候选|([中投]标|报价)(人|单位|候选)|单位名称|供应商)"
  316. height = len(inner_table)
  317. width = len(inner_table[0])
  318. text = ""
  319. for head_i in range(len(head_list)-1):
  320. head_begin = head_list[head_i]
  321. head_end = head_list[head_i+1]
  322. direct = getDirect(inner_table, head_begin, head_end)
  323. if direct=="row":
  324. for i in range(head_begin,head_end):
  325. rank_text = ""
  326. entity_text = ""
  327. text_line = ""
  328. #在同一句话中重复的可以去掉
  329. text_set = set()
  330. for j in range(width):
  331. cell = inner_table[i][j]
  332. #是属性值
  333. if cell[1]==0:
  334. find_flag = False
  335. head = ""
  336. temp_head = ""
  337. for loop_j in range(1,j+1):
  338. if inner_table[i][j-loop_j][1]==2:
  339. if find_flag:
  340. if inner_table[i][j-loop_j][0]!=temp_head:
  341. head = inner_table[i][j-loop_j][0]+":"+head
  342. else:
  343. head = inner_table[i][j-loop_j][0]+":"+head
  344. find_flag = True
  345. temp_head = inner_table[i][j-loop_j][0]
  346. else:
  347. if find_flag:
  348. break
  349. find_flag = False
  350. temp_head = ""
  351. for loop_i in range(0,i+1-head_begin):
  352. if inner_table[i-loop_i][j][1]==1:
  353. if find_flag:
  354. if inner_table[i-loop_i][j][0]!=temp_head:
  355. head = inner_table[i-loop_i][j][0]+":"+head
  356. else:
  357. head = inner_table[i-loop_i][j][0]+":"+head
  358. find_flag = True
  359. temp_head = inner_table[i-loop_i][j][0]
  360. else:
  361. #找到表头后遇到属性值就返回
  362. if find_flag:
  363. break
  364. if str(head+inner_table[i][j][0]) in text_set:
  365. continue
  366. if re.search(rankPattern,head) is not None:
  367. rank_text += head+inner_table[i][j][0]+","
  368. #print(rank_text)
  369. elif re.search(entityPattern,head) is not None:
  370. entity_text += head+inner_table[i][j][0]+","
  371. #print(entity_text)
  372. else:
  373. text_line += head+inner_table[i][j][0]+","
  374. text_set.add(str(head+inner_table[i][j][0]))
  375. text += rank_text+entity_text+text_line
  376. text = text[:-1]+"。"
  377. else:
  378. for j in range(width):
  379. rank_text = ""
  380. entity_text = ""
  381. text_line = ""
  382. text_set = set()
  383. for i in range(head_begin,head_end):
  384. cell = inner_table[i][j]
  385. #是属性值
  386. if cell[1]==0:
  387. find_flag = False
  388. head = ""
  389. temp_head = ""
  390. for loop_j in range(1,j+1):
  391. if inner_table[i][j-loop_j][1]==2:
  392. if find_flag:
  393. if inner_table[i][j-loop_j][0]!=temp_head:
  394. head = inner_table[i][j-loop_j][0]+":"+head
  395. else:
  396. head = inner_table[i][j-loop_j][0]+":"+head
  397. find_flag = True
  398. temp_head = inner_table[i][j-loop_j][0]
  399. else:
  400. if find_flag:
  401. break
  402. find_flag = False
  403. temp_head = ""
  404. for loop_i in range(0,i+1-head_begin):
  405. if inner_table[i-loop_i][j][1]==1:
  406. if find_flag:
  407. if inner_table[i-loop_i][j][0]!=temp_head:
  408. head = inner_table[i-loop_i][j][0]+":"+head
  409. else:
  410. head = inner_table[i-loop_i][j][0]+":"+head
  411. find_flag = True
  412. temp_head = inner_table[i-loop_i][j][0]
  413. else:
  414. if find_flag:
  415. break
  416. if str(head+inner_table[i][j][0]) in text_set:
  417. continue
  418. if re.search(rankPattern,head) is not None:
  419. rank_text += head+inner_table[i][j][0]+","
  420. #print(rank_text)
  421. elif re.search(entityPattern,head) is not None:
  422. entity_text += head+inner_table[i][j][0]+","
  423. #print(entity_text)
  424. else:
  425. text_line += head+inner_table[i][j][0]+","
  426. text_set.add(str(head+inner_table[i][j][0]))
  427. text += rank_text+entity_text+text_line
  428. text = text[:-1]+"。"
  429. return text
  430. def trunTable(tbody):
  431. fixSpan(tbody)
  432. inner_table = getTable(tbody)
  433. inner_table = fixTable(inner_table)
  434. if len(inner_table)>0 and len(inner_table[0])>0:
  435. #inner_table,head_list = setHead_withRule(inner_table,pat_head,pat_value,3)
  436. inner_table,head_list = setHead(inner_table)
  437. '''
  438. print("----")
  439. print(head_list)
  440. for item in inner_table:
  441. print(item)
  442. '''
  443. tbody.string = getTableText(inner_table,head_list)
  444. #print(tbody.string)
  445. tbody.name = "table"
  446. pat_head = re.compile('(名称|序号|项目|标项|工程|品目[一二三四1234]|第[一二三四1234](标段|名|候选人|中标)|包段|包号|货物|单位|数量|价格|报价|金额|总价|单价|[招投中]标|供应商|候选|编号|得分|评委|评分|名次|排名|排序|科室|方式|工期|时间|产品|开始|结束|联系|日期|面积|姓名|证号|备注|级别|地[点址]|类型|代理|制造)')
  447. #pat_head = re.compile('(名称|序号|项目|工程|品目[一二三四1234]|第[一二三四1234](标段|候选人|中标)|包段|包号|货物|单位|数量|价格|报价|金额|总价|单价|[招投中]标|供应商|候选|编号|得分|评委|评分|名次|排名|排序|科室|方式|工期|时间|产品|开始|结束|联系|日期|面积|姓名|证号|备注|级别|地[点址]|类型|代理)')
  448. pat_value = re.compile("(\d{2,}.\d{1}|\d+年\d+月|\d{8,}|\d{3,}-\d{6,}|有限[责任]*公司|^\d+$)")
  449. tbodies = soup.find_all('table')
  450. # 遍历表格中的每个tbody
  451. #逆序处理嵌套表格
  452. for tbody_index in range(1,len(tbodies)+1):
  453. tbody = tbodies[len(tbodies)-tbody_index]
  454. trunTable(tbody)
  455. tbodies = soup.find_all('tbody')
  456. # 遍历表格中的每个tbody
  457. #逆序处理嵌套表格
  458. for tbody_index in range(1,len(tbodies)+1):
  459. tbody = tbodies[len(tbodies)-tbody_index]
  460. trunTable(tbody)
  461. return soup
  462. def getSourceData():
  463. data = []
  464. data_set_is = set()
  465. data_set_no = set()
  466. for file in glob.glob("C:\\Users\\User\\Desktop\\20190320要素\\*.html"):
  467. filename = file.split("\\")[-1]
  468. source = codecs.open(file,"r",encoding="utf8").read()
  469. tableToText(BeautifulSoup(source,"lxml"),data,filename,data_set_is,data_set_no)
  470. for file in glob.glob("C:\\Users\\User\\Desktop\\20190306要素\\*.html"):
  471. filename = file.split("\\")[-1]
  472. source = codecs.open(file,"r",encoding="utf8").read()
  473. tableToText(BeautifulSoup(source,"lxml"),data,filename,data_set_is,data_set_no)
  474. ''''''
  475. list_file = []
  476. list_item = []
  477. list_label = []
  478. #data.sort(key=lambda x:x[2],reverse=True)
  479. data = data[0:60000]
  480. for item in data:
  481. list_file.append(item[0])
  482. list_item.append(item[1][:100])
  483. list_label.append(item[2])
  484. df = pd.DataFrame({"list_file":list_file,"list_item":list_item,"list_label":list_label})
  485. df.to_excel("data_item.xls",columns=["list_file","list_item","list_label"])
  486. def importData():
  487. conn = psycopg2.connect(dbname="article_label",user="postgres",password="postgres",host="192.168.2.101")
  488. cursor = conn.cursor()
  489. file = "data_item.xls"
  490. df = pd.read_excel(file)
  491. for file,text,label in zip(df["list_file"],df["list_item"],df["list_label"]):
  492. text = str(text)
  493. text = text.replace("\\","\\\\")
  494. text = re.sub("'","\\'",str(text))
  495. sql = " insert into form(filename,text,label) values(E'"+file+"',E'"+str(text)+"',E'"+str(int(label))+"')"
  496. print(sql)
  497. cursor.execute(sql)
  498. conn.commit()
  499. conn.close()
  500. def selectWithRule(source,filter,target):
  501. assert source!=target
  502. dict_source = pd.read_excel(source)
  503. set_filter = set()
  504. for filt in filter:
  505. set_filter = set_filter | set(pd.read_excel(filt)["list_item"])
  506. list_file = []
  507. list_item = []
  508. list_label = []
  509. for file,text,label in zip(dict_source["list_file"],dict_source["list_item"],dict_source["list_label"]):
  510. if str(text) in set_filter:
  511. continue
  512. if re.search(".{8,}(工程|项目|采购|公告|公示)",str(text)) is not None:
  513. #if len(str(text))>20:
  514. list_file.append(file)
  515. list_item.append(text)
  516. list_label.append(label)
  517. data = {"list_file":list_file,"list_item":list_item,"list_label":list_label}
  518. columns = ["list_file","list_item","list_label"]
  519. df = pd.DataFrame(data)
  520. df.to_excel(target,index=False,columns=columns)
  521. def importRelabel():
  522. files = ["批量.xls"]
  523. conn = psycopg2.connect(dbname="article_label",user="postgres",password="postgres",host="192.168.2.101")
  524. cursor = conn.cursor()
  525. for file in files:
  526. df = pd.read_excel(file)
  527. for text,relabel in zip(df["list_item"],df["list_relabel"]):
  528. text = str(text)
  529. text = text.replace("\\","\\\\")
  530. text = re.sub("'","\\'",str(text))
  531. sql = " update form set relabel='"+str(int(relabel))+"' where text=E'"+str(text)+"' "
  532. cursor.execute(sql)
  533. conn.commit()
  534. conn.close()
  535. def getHtml():
  536. conn = psycopg2.connect(dbname="article_label",user="postgres",password="postgres",host="192.168.2.101")
  537. cursor = conn.cursor()
  538. sql = " select filename from form where relabel is NULL group by filename having count(1)>0 "
  539. cursor.execute(sql)
  540. rows = cursor.fetchall()
  541. data = []
  542. index = 0
  543. for row in rows:
  544. filename = row[0]
  545. if filename=="比地_101_58519594.html":
  546. print(index)
  547. path = "C:\\Users\\User\\Desktop\\20190320要素\\"+filename
  548. if not os.path.exists(path):
  549. path = "C:\\Users\\User\\Desktop\\20190306要素\\"+filename
  550. data.append([filename,codecs.open(path,'r',encoding="utf8").read()])
  551. index += 1
  552. #save(data,"namehtml.pk")
  553. def getTrainData(percent=0.9):
  554. conn = psycopg2.connect(dbname="article_label",user="postgres",password="postgres",host="192.168.2.101")
  555. cursor = conn.cursor()
  556. sql = "select filename,text,label,relabel,handlabel from form "
  557. cursor.execute(sql)
  558. rows = cursor.fetchall()
  559. save(rows,"filename_text_label_relabel_handlabel.pk")
  560. train_x = []
  561. train_y = []
  562. test_x = []
  563. test_y = []
  564. test_text = []
  565. for row in rows:
  566. input = str(row[1])
  567. label = str(int(row[2]))
  568. if row[4] is not None:
  569. label = str(int(row[4]))
  570. elif row[3] is not None:
  571. label = str(int(row[3]))
  572. item_y = [0,0]
  573. item_y[int(label)] = 1
  574. if np.random.random()<percent:
  575. # train_x.append(encodeInput(input))
  576. train_x.append(encodeInput([input], word_len=50, word_flag=True,userFool=False)[0])
  577. train_y.append(item_y)
  578. else:
  579. # test_x.append(encodeInput(input))
  580. test_x.append(encodeInput([input], word_len=50, word_flag=True,userFool=False)[0])
  581. test_y.append(item_y)
  582. test_text.append([row[0],input])
  583. return np.array(train_x),np.array(train_y),np.array(test_x),np.array(test_y),test_text
  584. def getTrainData_jsonTable(begin,end,return_text=False):
  585. def encode_table(inner_table,size=30):
  586. def encode_item(_table,i,j):
  587. _x = [_table[j-1][i-1],_table[j-1][i],_table[j-1][i+1],
  588. _table[j][i-1],_table[j][i],_table[j][i+1],
  589. _table[j+1][i-1],_table[j+1][i],_table[j+1][i+1]]
  590. e_x = [encodeInput_form(_temp[0],MAX_LEN=30) for _temp in _x]
  591. _label = _table[j][i][1]
  592. # print(_x)
  593. # print(_x[4],_label)
  594. return e_x,_label,_x
  595. def copytable(inner_table):
  596. table = []
  597. for line in inner_table:
  598. list_line = []
  599. for item in line:
  600. list_line.append([item[0][:size],item[1]])
  601. table.append(list_line)
  602. return table
  603. table = copytable(inner_table)
  604. padding = ["#"*30,0]
  605. width = len(table[0])
  606. height = len(table)
  607. table.insert(0,[padding for i in range(width)])
  608. table.append([padding for i in range(width)])
  609. for item in table:
  610. item.insert(0,padding.copy())
  611. item.append(padding.copy())
  612. data_x = []
  613. data_y = []
  614. data_text = []
  615. data_position = []
  616. for _i in range(1,width+1):
  617. for _j in range(1,height+1):
  618. _x,_y,_text = encode_item(table,_i,_j)
  619. data_x.append(_x)
  620. _label = [0,0]
  621. _label[_y] = 1
  622. data_y.append(_label)
  623. data_text.append(_text)
  624. data_position.append([_i-1,_j-1])
  625. # input = table[_j][_i][0]
  626. # item_y = [0,0]
  627. # item_y[table[_j][_i][1]] = 1
  628. # data_x.append(encodeInput([input], word_len=50, word_flag=True,userFool=False)[0])
  629. # data_y.append(item_y)
  630. return data_x,data_y,data_text,data_position
  631. def getDataSet(list_json_table,return_text=False):
  632. _count = 0
  633. _sum = len(list_json_table)
  634. data_x = []
  635. data_y = []
  636. data_text = []
  637. for json_table in list_json_table:
  638. _count += 1
  639. print("%d/%d"%(_count,_sum))
  640. table = json.loads(json_table)
  641. if table is not None:
  642. list_x,list_y,list_text = encode_table(table)
  643. data_x.extend(list_x)
  644. data_y.extend(list_y)
  645. if return_text:
  646. data_text.extend(list_text)
  647. return np.array(data_x),np.array(data_y),data_text
  648. save_path = "./traindata/websource_67000_table_%d-%d-%s.pk"%(begin,end,"1" if return_text else "0")
  649. if os.path.exists(save_path):
  650. data_x,data_y,data_text = load(save_path)
  651. else:
  652. df = pd.read_csv("../../dl_dev/form/traindata/websource_67000_table.csv", encoding="GBK")
  653. import json
  654. data_x,data_y,data_text = getDataSet(df["json_table"][begin:end],return_text=return_text)
  655. save((data_x,data_y,data_text),save_path)
  656. return data_x,data_y,data_text
  657. if __name__=="__main__":
  658. #getSourceData()
  659. #importData()
  660. #selectWithRule("data_item.xls", ["批量.xls"], "temp.xls")
  661. #importRelabel()
  662. # getHtml()
  663. getTrainData_jsonTable()