check_user_label_accuracy.py 6.4 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199
  1. import time
  2. from BiddingKG.dl.table_head.pre_process import postgresql_util
  3. from BiddingKG.dl.table_head.predict import predict
  4. def user_label_accuracy(update_user):
  5. if update_user != 'test':
  6. sql = """
  7. select table_text, pre_label, post_label, id
  8. from label_table_head_info
  9. where update_user='""" + update_user + "' order by update_time"
  10. else:
  11. sql = """
  12. select table_text, pre_label, post_label, id
  13. from label_table_head_info
  14. where update_user='""" + update_user + "' and status = 1 and update_time >= '2022-01-23'"
  15. result_list = postgresql_util(sql, limit=1000000)
  16. right_cnt = 0
  17. error_cnt = 0
  18. error_id_list = []
  19. right_id_list = []
  20. i = 0
  21. start_time = time.time()
  22. for table in result_list:
  23. i += 1
  24. if i % 1000 == 0:
  25. if right_cnt + error_cnt != 0:
  26. print("Loop", i, right_cnt/(right_cnt+error_cnt), time.time()-start_time)
  27. else:
  28. print("Loop", i, time.time()-start_time)
  29. start_time = time.time()
  30. pre_label = eval(table[1])
  31. post_label = eval(table[2])
  32. _id = table[3]
  33. # table_text需要特殊处理
  34. try:
  35. table_text = table[0]
  36. if table_text[0] == '"':
  37. table_text = eval(table_text)
  38. else:
  39. table_text = table_text
  40. table_text = table_text.replace('\\', '/')
  41. table_text = eval(table_text)
  42. except:
  43. print("无法识别table_text", _id)
  44. continue
  45. if post_label:
  46. label_list = post_label
  47. else:
  48. label_list = pre_label
  49. predict_label_list = predict(table_text, model_id=3)
  50. if predict_label_list:
  51. if str(label_list) == str(predict_label_list):
  52. right_id_list.append(str(_id)+"\n")
  53. # right_cnt += 1
  54. else:
  55. error_id_list.append(str(_id)+"\n")
  56. # error_cnt += 1
  57. if len(label_list) == len(predict_label_list):
  58. for j in range(len(label_list)):
  59. for k in range(len(label_list[j])):
  60. if table_text[j][k] == "":
  61. continue
  62. if label_list[j][k] == "1" or predict_label_list[j][k] == "1":
  63. if len(table_text[j][k]) >= 20:
  64. continue
  65. else:
  66. if label_list[j][k] == predict_label_list[j][k]:
  67. right_cnt += 1
  68. else:
  69. error_cnt += 1
  70. else:
  71. print("len(label_list) == len(predict_label_list)", _id,
  72. len(label_list), len(predict_label_list))
  73. accuracy = right_cnt / (right_cnt + error_cnt)
  74. print(update_user + " accuracy:", accuracy, 'total:', len(result_list))
  75. print("error_id_list", len(error_id_list))
  76. save_path = "check_user_result/accuracy.txt"
  77. with open(save_path, 'a') as f:
  78. f.write(update_user + " "
  79. + "表头正确率-" + str(round(accuracy, 2)) + " "
  80. + "文章数-" + str(len(result_list)) + " "
  81. + "表格总数-" + str(right_cnt + error_cnt) + " "
  82. + "表头正确数-" + str(right_cnt)
  83. + "\n")
  84. save_path = "check_user_result/"+update_user+"_error.txt"
  85. with open(save_path, 'w') as f:
  86. f.writelines(error_id_list)
  87. save_path = "check_user_result/"+update_user+"_right.txt"
  88. with open(save_path, 'w') as f:
  89. f.writelines(right_id_list)
  90. return accuracy
  91. def get_single_result(_id):
  92. sql = """
  93. select table_text, pre_label, post_label, id
  94. from label_table_head_info
  95. where id=""" + str(_id)
  96. result_list = postgresql_util(sql, limit=1000000)
  97. table_text = result_list[0][0]
  98. if table_text[0] == '"':
  99. table_text = eval(table_text)
  100. else:
  101. table_text = table_text
  102. table_text = table_text.replace('\\', '/')
  103. table_text = eval(table_text)
  104. label_list = predict(table_text)
  105. for i in range(len(label_list)):
  106. print(i+1, label_list[i])
  107. if __name__ == '__main__':
  108. # users = ["test9", "test11", "test12", "test25", "test26"]
  109. # users = ["test9", "test11", ]
  110. # users = ['test12', 'test25']
  111. # users = ["test20", "test27"]
  112. # users = ['test']
  113. users = [
  114. "test1",
  115. "test11",
  116. "test12",
  117. "test16",
  118. "test17",
  119. "test19",
  120. "test20",
  121. "test21",
  122. "test22",
  123. "test25",
  124. "test26",
  125. "test27",
  126. "test29",
  127. "test3",
  128. "test7",
  129. "test8",
  130. "test9",
  131. ]
  132. users = ["test"]
  133. users = ["test12", "test17", "test21", "test22", "test27", ]
  134. users = ["test27"]
  135. acc_list = []
  136. for user in users:
  137. acc = user_label_accuracy(user)
  138. acc_list.append([user, acc])
  139. print(acc_list)
  140. # get_single_result(161927)
  141. # import pandas as pd
  142. # df = pd.read_csv("C:\\Users\\Administrator\\Desktop\\4.csv")
  143. # _dict = {
  144. # 51: "公告变更",
  145. # 52: "招标公告",
  146. # 101: "中标信息",
  147. # 102: "招标预告",
  148. # 103: "招标答疑",
  149. # 104: "招标文件",
  150. # 105: "资审结果",
  151. # 106: "法律法规",
  152. # 107: "新闻资讯",
  153. # 108: "拟建项目",
  154. # 109: "展会推广",
  155. # 110: "企业名录",
  156. # 111: "企业资质",
  157. # 112: "全国工程人员",
  158. # 113: "业主采购",
  159. # 114: "采购意向",
  160. # 115: "拍卖出让",
  161. # 116: "土地矿产",
  162. # 117: "产权交易",
  163. # 118: "废标公告",
  164. # 119: "候选人公示",
  165. # 120: "合同公告",
  166. # }
  167. # data_list = []
  168. # for index, row in df.iterrows():
  169. # if index % 100000 == 0:
  170. # print("Loop", index)
  171. # print(_dict[int(row['docchannel'])])
  172. # print(df.iloc[index, 2])
  173. # data = row.tolist()
  174. # data[2] = _dict[int(data[2])]
  175. # data_list.append(data)
  176. # df = pd.DataFrame(data_list)
  177. # df.columns = ['docid', '项目名称', '信息类型', '发布时间', '地区', '业主',
  178. # '预算金额', '中标供应商', '成交金额', '代理机构']
  179. # df.to_csv("C:\\Users\\Administrator\\Desktop\\4-1.csv", index=False)