pre_process.py 19 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302303304305306307308309310311312313314315316317318319320321322323324325326327328329330331332333334335336337338339340341342343344345346347348349350351352353354355356357358359360361362363364365366367368369370371372373374375376377378379380381382383384385386387388389390391392393394395396397398399400401402403404405406407408409410411412413414415416417418419420421422423424425426427428429430431432433434435436437438439440441442443444445446447448449450451452453454455456457458459460461462463464465466467468469470471472473474475476477478479480481482483484485486487488489490491492493494495496497498499500501502503504505506507508509510511512513514515516517518519520521522523524525526527528529530531532533534535536537538539540541542543544545546547548549550551552553554555556557558559560561562563564565566567568569570571572573574575576577578579580581582583584585586587588589590591592593594595
  1. # DEPRECATED(Phase1): 本文件中的 PostgreSQL 硬编码连接(host=192.168.*,
  2. # password=postgres 等)将在后续 training/ Phase 迁移到 BiddingKG.dl.infra.db。
  3. # 迁移完成前可临时使用:from BiddingKG.dl.infra.db import get_connection;
  4. # conn = get_connection("<dbname>")
  5. # 详见 ARCHITECTURE.md 第 11 章 Phase 1 与 REFACTOR_LOG.md。
  6. import os
  7. import random
  8. import sys
  9. import numpy as np
  10. sys.path.append(os.path.dirname(__file__) + "/../")
  11. from common.Utils import embedding_word, embedding_word_forward
  12. def get_sentence_index_list(sentence, dict_path='utils/ppocr_keys_v1.txt'):
  13. with open(dict_path, 'r') as f:
  14. character_list = f.readlines()
  15. for i in range(len(character_list)):
  16. character_list[i] = character_list[i][:-1]
  17. index_list = []
  18. for character in sentence:
  19. if character == '':
  20. index_list.append(0)
  21. elif character in character_list:
  22. _index = character_list.index(character) + 1
  23. index_list.append(_index)
  24. else:
  25. index_list.append(0)
  26. return index_list
  27. def postgresql_util(sql, limit):
  28. import psycopg2
  29. conn = psycopg2.connect(dbname="table_head_label", user="postgres", password="postgres",
  30. host="192.168.2.103")
  31. cursor = conn.cursor()
  32. cursor.execute(sql)
  33. print(sql)
  34. rows = cursor.fetchmany(1000)
  35. cnt = 0
  36. all_rows = []
  37. while rows:
  38. if cnt >= limit:
  39. break
  40. all_rows += rows
  41. cnt += len(rows)
  42. rows = cursor.fetchmany(1000)
  43. return all_rows
  44. def get_data_from_sql(dim=10, whole_table=False, padding=True):
  45. sql = """
  46. select table_text, pre_label, post_label, id
  47. from label_table_head_info
  48. where status = 0 and (update_user='test9' or update_user='test1' or update_user='test7' or update_user='test26')
  49. ;
  50. """
  51. # sql = """
  52. # select table_text, pre_label, post_label, id
  53. # from label_table_head_info
  54. # where status = 1 and update_time >= '2022-01-17' and update_time <= '2022-01-22'
  55. # ;
  56. # """
  57. result_list = postgresql_util(sql, limit=1000000)
  58. # 需排除的id
  59. with open(r"C:\Users\Administrator\Desktop\table_not_eval.txt", "r") as f:
  60. delete_id_list = eval(f.read())
  61. with open(r"C:\Users\Administrator\Desktop\table_delete.txt", "r") as f:
  62. delete_id_list += eval(f.read())
  63. all_data_list = []
  64. all_data_label_list = []
  65. i = 0
  66. # 一行就是一篇表格
  67. for table in result_list:
  68. i += 1
  69. if i % 100 == 0:
  70. print("Loop", i)
  71. pre_label = eval(table[1])
  72. post_label = eval(table[2])
  73. _id = table[3]
  74. if _id in delete_id_list:
  75. print("pass", _id)
  76. continue
  77. # table_text需要特殊处理
  78. try:
  79. table_text = table[0]
  80. if table_text[0] == '"':
  81. table_text = eval(table_text)
  82. else:
  83. table_text = table_text
  84. table_text = table_text.replace('\\', '/')
  85. table_text = eval(table_text)
  86. except:
  87. print("无法识别table_text", _id)
  88. continue
  89. if whole_table:
  90. if len(post_label) >= 2:
  91. data_list, data_label_list = table_pre_process_2(table_text, post_label,
  92. _id, padding=padding)
  93. elif len(pre_label) >= 2:
  94. data_list, data_label_list = table_pre_process_2(table_text, pre_label,
  95. _id, padding=padding)
  96. else:
  97. data_list, data_label_list = [], []
  98. else:
  99. # 只有一行的也不要
  100. if len(post_label) >= 2:
  101. data_list, data_label_list = table_pre_process(table_text, post_label, _id)
  102. elif len(pre_label) >= 2:
  103. data_list, data_label_list = table_pre_process(table_text, pre_label, _id)
  104. else:
  105. data_list, data_label_list = [], []
  106. all_data_list += data_list
  107. all_data_label_list += data_label_list
  108. # 按维度大小排序
  109. if whole_table:
  110. _list = []
  111. for data, label in zip(all_data_list, all_data_label_list):
  112. _list.append([data, label])
  113. _list.sort(key=lambda x: (len(x[0]), len(x[0][0])))
  114. all_data_list[:], all_data_label_list[:] = zip(*_list)
  115. print("len(all_data_list)", len(all_data_list))
  116. return all_data_list, all_data_label_list
  117. def table_pre_process(text_list, label_list, _id, is_train=True):
  118. """
  119. 表格处理,每个单元格生成2条数据,横竖各1条
  120. :param text_list:
  121. :param label_list:
  122. :param _id:
  123. :param is_train:
  124. :return:
  125. """
  126. if is_train:
  127. if len(text_list) != len(label_list):
  128. print("文字单元格与标注单元格数量不匹配!", _id)
  129. print("len(text_list)", len(text_list), "len(label_list)", len(label_list))
  130. return [], []
  131. data_list = []
  132. data_label_list = []
  133. for i in range(len(text_list)):
  134. row = text_list[i]
  135. if is_train:
  136. row_label = label_list[i]
  137. if i > 0:
  138. last_row = text_list[i-1]
  139. if is_train:
  140. last_row_label = label_list[i-1]
  141. else:
  142. last_row = []
  143. if is_train:
  144. last_row_label = []
  145. if i < len(text_list) - 1:
  146. next_row = text_list[i+1]
  147. if is_train:
  148. next_row_label = label_list[i+1]
  149. else:
  150. next_row = []
  151. if is_train:
  152. next_row_label = []
  153. for j in range(len(row)):
  154. col = row[j]
  155. if is_train:
  156. col_label = row_label[j]
  157. # 超出表格置为None, 0
  158. if j > 0:
  159. last_col = row[j-1]
  160. if is_train:
  161. last_col_label = row_label[j-1]
  162. else:
  163. last_col = col
  164. if is_train:
  165. last_col_label = col_label
  166. if j < len(row) - 1:
  167. next_col = row[j+1]
  168. if is_train:
  169. next_col_label = row_label[j+1]
  170. else:
  171. next_col = col
  172. if is_train:
  173. next_col_label = col_label
  174. if last_row:
  175. last_row_col = last_row[j]
  176. if is_train:
  177. last_row_col_label = last_row_label[j]
  178. else:
  179. last_row_col = col
  180. if is_train:
  181. last_row_col_label = col_label
  182. if next_row:
  183. next_row_col = next_row[j]
  184. if is_train:
  185. next_row_col_label = next_row_label[j]
  186. else:
  187. next_row_col = col
  188. if is_train:
  189. next_row_col_label = col_label
  190. # data_list.append([last_col, col, next_col])
  191. # if is_train:
  192. # data_label_list.append([int(last_col_label), int(col_label),
  193. # int(next_col_label)])
  194. #
  195. # data_list.append([last_row_col, col, next_row_col])
  196. # if is_train:
  197. # data_label_list.append([int(last_row_col_label), int(col_label),
  198. # int(next_row_col_label)])
  199. if is_train:
  200. dup_list = [str(x) for x in data_list]
  201. data = [last_col, col, next_col, last_row_col, col, next_row_col]
  202. if str(data) not in dup_list:
  203. data_list.append([last_col, col, next_col, last_row_col, col, next_row_col])
  204. data_label_list.append(int(col_label))
  205. else:
  206. data_list.append([last_col, col, next_col, last_row_col, col, next_row_col])
  207. if is_train:
  208. return data_list, data_label_list
  209. else:
  210. return data_list
  211. def table_pre_process_2(text_list, label_list, _id, is_train=True, padding=True):
  212. """
  213. 表格处理,整个表格为一个数组,且填充长宽维度
  214. :param text_list:
  215. :param label_list:
  216. :param _id:
  217. :param is_train:
  218. :return:
  219. """
  220. # 判断表格长宽是否合理
  221. row_len = len(text_list)
  222. best_row_len = get_best_padding_size(row_len, min_len=8)
  223. col_len = len(text_list[0])
  224. best_col_len = get_best_padding_size(col_len, min_len=8)
  225. if best_row_len is None:
  226. if is_train:
  227. return [], []
  228. else:
  229. return []
  230. if best_col_len is None:
  231. if is_train:
  232. return [], []
  233. else:
  234. return []
  235. if is_train:
  236. if len(text_list) != len(label_list):
  237. print("文字单元格与标注单元格数量不匹配!", _id)
  238. print("len(text_list)", len(text_list), "len(label_list)", len(label_list))
  239. return [], []
  240. if padding:
  241. for i in range(row_len):
  242. col_len = len(text_list[i])
  243. text_list[i] += [None]*(best_col_len-col_len)
  244. if is_train:
  245. label_list[i] += ["0"]*(best_col_len-col_len)
  246. text_list += [[None]*best_col_len]*(best_row_len-row_len)
  247. if is_train:
  248. label_list += [["0"]*best_col_len]*(best_row_len-row_len)
  249. if is_train:
  250. for i in range(len(label_list)):
  251. for j in range(len(label_list[i])):
  252. label_list[i][j] = int(label_list[i][j])
  253. return [text_list], [label_list]
  254. else:
  255. return [text_list]
  256. def get_best_padding_size(axis_len, min_len=3, max_len=300):
  257. # sizes = [8, 16, 24, 32, 40, 48, 56, 64, 72, 80, 88, 96, 104, 112, 120,
  258. # 128, 136, 144, 152, 160, 168, 176, 184, 192, 200, 208, 216, 224,
  259. # 232, 240, 248, 256, 264, 272, 280, 288, 296]
  260. # sizes = [3, 6, 9, 12, 15, 18, 21, 24, 27, 30, 33, 36, 39, 42, 45, 48, 51, 54, 57,
  261. # 60, 63, 66, 69, 72, 75, 78, 81, 84, 87, 90, 93, 96, 99, 102, 105, 108, 111,
  262. # 114, 117, 120, 123, 126, 129, 132, 135, 138, 141, 144, 147, 150, 153, 156,
  263. # 159, 162, 165, 168, 171, 174, 177, 180, 183, 186, 189, 192, 195, 198, 201,
  264. # 204, 207, 210, 213, 216, 219, 222, 225, 228, 231, 234, 237, 240, 243, 246,
  265. # 249, 252, 255, 258, 261, 264, 267, 270, 273, 276, 279, 282, 285, 288, 291,
  266. # 294, 297]
  267. sizes = []
  268. for i in range(1, max_len):
  269. if i * min_len <= max_len:
  270. sizes.append(i * min_len)
  271. if axis_len > sizes[-1]:
  272. return axis_len
  273. best_len = sizes[-1]
  274. for height in sizes:
  275. if axis_len <= height:
  276. best_len = height
  277. break
  278. # print("get_best_padding_size", axis_len, best_len)
  279. return best_len
  280. def get_data_from_file(file_type, model_id=1):
  281. if file_type == 'np':
  282. data_path = 'train_data/data_3.npy'
  283. data_label_path = 'train_data/data_label_3.npy'
  284. array1 = np.load(data_path)
  285. array2 = np.load(data_label_path)
  286. return array1, array2
  287. elif file_type == 'txt':
  288. if model_id == 1:
  289. data_path = 'train_data/data1.txt'
  290. data_label_path = 'train_data/data_label1.txt'
  291. elif model_id == 2:
  292. data_path = 'train_data/data2.txt'
  293. data_label_path = 'train_data/data_label2.txt'
  294. elif model_id == 3:
  295. data_path = 'train_data/data3.txt'
  296. data_label_path = 'train_data/data_label3.txt'
  297. with open(data_path, 'r') as f:
  298. data_list = f.readlines()
  299. with open(data_label_path, 'r') as f:
  300. data_label_list = f.readlines()
  301. return data_list, data_label_list
  302. else:
  303. print("file type error! only np and txt supported")
  304. raise Exception
  305. def processed_save_to_np():
  306. array1, array2 = get_data_from_sql()
  307. np.save('train_data/data_3.npy', array1)
  308. np.save('train_data/data_label_3.npy', array2)
  309. # with open('train_data/data.txt', 'w') as f:
  310. # for line in list1:
  311. # f.write(str(line) + "\n")
  312. # with open('train_data/data_label.txt', 'w') as f:
  313. # for line in list2:
  314. # f.write(str(line) + "\n")
  315. def processed_save_to_txt(whole_table=False, padding=True):
  316. list1, list2 = get_data_from_sql(whole_table=whole_table, padding=padding)
  317. # 打乱
  318. # if not whole_table or not padding:
  319. zip_list = list(zip(list1, list2))
  320. random.shuffle(zip_list)
  321. list1[:], list2[:] = zip(*zip_list)
  322. with open('train_data/data1.txt', 'w') as f:
  323. for line in list1:
  324. f.write(str(line) + "\n")
  325. with open('train_data/data_label1.txt', 'w') as f:
  326. for line in list2:
  327. f.write(str(line) + "\n")
  328. def data_balance():
  329. data_list, data_label_list = get_data_from_file('txt')
  330. all_cnt = len(data_label_list)
  331. cnt_0 = 0
  332. cnt_1 = 0
  333. for data in data_label_list:
  334. if eval(data[:-1])[1] == 1:
  335. cnt_1 += 1
  336. else:
  337. cnt_0 += 1
  338. print("all_cnt", all_cnt)
  339. print("label has 1", cnt_1)
  340. print("label all 0", cnt_0)
  341. def test_embedding():
  342. output_shape = (2, 1, 60)
  343. data = [[None], [None]]
  344. result = embedding_word(data, output_shape)
  345. print(result)
  346. def my_data_loader(data_list, data_label_list, batch_size, is_train=True):
  347. data_num = len(data_list)
  348. # 定义Embedding输出
  349. output_shape = (6, 20, 60)
  350. # batch循环取数据
  351. i = 0
  352. if is_train:
  353. while True:
  354. new_data_list = []
  355. new_data_label_list = []
  356. for j in range(batch_size):
  357. if i >= data_num:
  358. i = 0
  359. # 中文字符映射为Embedding
  360. data = eval(data_list[i][:-1])
  361. data_label = eval(data_label_list[i][:-1])
  362. data = embedding_word(data, output_shape)
  363. if data.shape == output_shape:
  364. new_data_list.append(data)
  365. new_data_label_list.append(data_label)
  366. i += 1
  367. new_data_list = np.array(new_data_list)
  368. new_data_label_list = np.array(new_data_label_list)
  369. X = new_data_list
  370. Y = new_data_label_list
  371. # (table_num, 3 sentences, dim characters, embedding) -> (3, table_num, dim, embedding)
  372. X = np.transpose(X, (1, 0, 2, 3))
  373. if (X[0] == X[1]).all():
  374. X[0] = np.zeros_like(X[1], dtype='float32')
  375. if (X[2] == X[1]).all():
  376. X[2] = np.zeros_like(X[1], dtype='float32')
  377. if (X[3] == X[1]).all():
  378. X[3] = np.zeros_like(X[1], dtype='float32')
  379. if (X[5] == X[1]).all():
  380. X[5] = np.zeros_like(X[1], dtype='float32')
  381. yield {'input_1': X[0], 'input_2': X[1], 'input_3': X[2],
  382. 'input_4': X[3], 'input_5': X[4], 'input_6': X[5]}, \
  383. {'output': Y}
  384. else:
  385. new_data_list = []
  386. for j in range(len(data_list)):
  387. # 中文字符映射为Embedding
  388. data = data_list[i]
  389. data = embedding_word(data, output_shape)
  390. if data.shape == output_shape:
  391. new_data_list.append(data)
  392. i += 1
  393. for j in range(0, len(data_list), batch_size):
  394. sub_data_list = np.array(new_data_list[j: j+batch_size])
  395. X = sub_data_list
  396. X = np.transpose(X, (1, 0, 2, 3))
  397. # print(X)
  398. # return X
  399. yield {'input_1': X[0], 'input_2': X[1], 'input_3': X[2],
  400. 'input_4': X[3], 'input_5': X[4], 'input_6': X[5], }
  401. def my_data_loader_predict(data_list, data_label_list, batch_size):
  402. data_num = len(data_list)
  403. # 定义Embedding输出
  404. output_shape = (6, 20, 60)
  405. i = 0
  406. new_data_list = []
  407. for j in range(len(data_list)):
  408. # 中文字符映射为Embedding
  409. data = data_list[i]
  410. data = embedding_word(data, output_shape)
  411. if data.shape == output_shape:
  412. new_data_list.append(data)
  413. i += 1
  414. sub_data_list = np.array(new_data_list)
  415. X = sub_data_list
  416. X = np.transpose(X, (1, 0, 2, 3))
  417. return X
  418. def my_data_loader_2(table_list, table_label_list, batch_size, is_train=True):
  419. pad_len = 0
  420. table_num = len(table_list)
  421. if is_train and batch_size == 1:
  422. table_list, table_label_list = get_random(table_list, table_label_list)
  423. # Embedding shape
  424. output_shape = (20, 60)
  425. # batch循环取数据
  426. i = 0
  427. last_shape = None
  428. while True:
  429. new_table_list = []
  430. new_table_label_list = []
  431. for j in range(batch_size):
  432. if i >= table_num:
  433. i = 0
  434. if is_train:
  435. table_list, table_label_list = get_random(table_list, table_label_list,
  436. seed=random.randint(1, 40))
  437. if type(table_list[i]) != list:
  438. table = eval(table_list[i][:-1])
  439. else:
  440. table = table_list[i]
  441. if batch_size > 1:
  442. if last_shape is None:
  443. last_shape = (len(table), len(table[0]))
  444. continue
  445. if (len(table), len(table[0])) != last_shape:
  446. last_shape = (len(table), len(table[0]))
  447. break
  448. if is_train:
  449. table_label = eval(table_label_list[i][:-1])
  450. # 中文字符映射为Embedding
  451. for k in range(len(table)):
  452. table[k] = embedding_word_forward(table[k], (len(table[k]),
  453. output_shape[0],
  454. output_shape[1]))
  455. new_table_list.append(table)
  456. if is_train:
  457. new_table_label_list.append(table_label)
  458. i += 1
  459. new_table_list = np.array(new_table_list)
  460. X = new_table_list
  461. if X.shape[-2:] != output_shape:
  462. # print("Dimension not match!", X.shape)
  463. # print("\n")
  464. continue
  465. # 获取Padding大小
  466. pad_height = get_best_padding_size(X.shape[1], pad_len)
  467. pad_width = get_best_padding_size(X.shape[2], pad_len)
  468. input_2 = np.zeros([1, X.shape[1], X.shape[2], pad_height, pad_width])
  469. if is_train:
  470. new_table_label_list = np.array(new_table_label_list)
  471. Y = new_table_label_list
  472. # Y = Y.astype(np.float32)
  473. # yield {"input_1": X, "input_2": input_2}, \
  474. # {"output_1": Y, "output_2": Y}
  475. yield {"input_1": X, "input_2": input_2}, \
  476. {"output": Y}
  477. else:
  478. yield {"input_1": X, "input_2": input_2}
  479. def check_train_data():
  480. data_list, label_list = get_data_from_file('txt', model_id=2)
  481. for data in data_list:
  482. data = eval(data)
  483. if len(data) % 8 != 0:
  484. print(len(data))
  485. print(len(data[0]))
  486. for row in data:
  487. if len(row) % 8 != 0:
  488. print(len(data))
  489. print(len(row))
  490. def get_random(text_list, label_list, seed=42):
  491. random.seed(seed)
  492. zip_list = list(zip(text_list, label_list))
  493. random.shuffle(zip_list)
  494. text_list[:], label_list[:] = zip(*zip_list)
  495. return text_list, label_list
  496. if __name__ == '__main__':
  497. processed_save_to_txt(whole_table=False, padding=False)
  498. # data_balance()
  499. # test_embedding()
  500. # check_train_data()
  501. # _list = []
  502. # for i in range(1, 100):
  503. # _list.append(i*3)
  504. # print(_list)
  505. # print(get_best_padding_size(9, 5))