| 123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199 |
- import time
- from BiddingKG.dl.table_head.pre_process import postgresql_util
- from BiddingKG.dl.table_head.predict import predict
- def user_label_accuracy(update_user):
- if update_user != 'test':
- sql = """
- select table_text, pre_label, post_label, id
- from label_table_head_info
- where update_user='""" + update_user + "' order by update_time"
- else:
- sql = """
- select table_text, pre_label, post_label, id
- from label_table_head_info
- where update_user='""" + update_user + "' and status = 1 and update_time >= '2022-01-23'"
- result_list = postgresql_util(sql, limit=1000000)
- right_cnt = 0
- error_cnt = 0
- error_id_list = []
- right_id_list = []
- i = 0
- start_time = time.time()
- for table in result_list:
- i += 1
- if i % 1000 == 0:
- if right_cnt + error_cnt != 0:
- print("Loop", i, right_cnt/(right_cnt+error_cnt), time.time()-start_time)
- else:
- print("Loop", i, time.time()-start_time)
- start_time = time.time()
- pre_label = eval(table[1])
- post_label = eval(table[2])
- _id = table[3]
- # table_text需要特殊处理
- try:
- table_text = table[0]
- if table_text[0] == '"':
- table_text = eval(table_text)
- else:
- table_text = table_text
- table_text = table_text.replace('\\', '/')
- table_text = eval(table_text)
- except:
- print("无法识别table_text", _id)
- continue
- if post_label:
- label_list = post_label
- else:
- label_list = pre_label
- predict_label_list = predict(table_text, model_id=3)
- if predict_label_list:
- if str(label_list) == str(predict_label_list):
- right_id_list.append(str(_id)+"\n")
- # right_cnt += 1
- else:
- error_id_list.append(str(_id)+"\n")
- # error_cnt += 1
- if len(label_list) == len(predict_label_list):
- for j in range(len(label_list)):
- for k in range(len(label_list[j])):
- if table_text[j][k] == "":
- continue
- if label_list[j][k] == "1" or predict_label_list[j][k] == "1":
- if len(table_text[j][k]) >= 20:
- continue
- else:
- if label_list[j][k] == predict_label_list[j][k]:
- right_cnt += 1
- else:
- error_cnt += 1
- else:
- print("len(label_list) == len(predict_label_list)", _id,
- len(label_list), len(predict_label_list))
- accuracy = right_cnt / (right_cnt + error_cnt)
- print(update_user + " accuracy:", accuracy, 'total:', len(result_list))
- print("error_id_list", len(error_id_list))
- save_path = "check_user_result/accuracy.txt"
- with open(save_path, 'a') as f:
- f.write(update_user + " "
- + "表头正确率-" + str(round(accuracy, 2)) + " "
- + "文章数-" + str(len(result_list)) + " "
- + "表格总数-" + str(right_cnt + error_cnt) + " "
- + "表头正确数-" + str(right_cnt)
- + "\n")
- save_path = "check_user_result/"+update_user+"_error.txt"
- with open(save_path, 'w') as f:
- f.writelines(error_id_list)
- save_path = "check_user_result/"+update_user+"_right.txt"
- with open(save_path, 'w') as f:
- f.writelines(right_id_list)
- return accuracy
- def get_single_result(_id):
- sql = """
- select table_text, pre_label, post_label, id
- from label_table_head_info
- where id=""" + str(_id)
- result_list = postgresql_util(sql, limit=1000000)
- table_text = result_list[0][0]
- if table_text[0] == '"':
- table_text = eval(table_text)
- else:
- table_text = table_text
- table_text = table_text.replace('\\', '/')
- table_text = eval(table_text)
- label_list = predict(table_text)
- for i in range(len(label_list)):
- print(i+1, label_list[i])
- if __name__ == '__main__':
- # users = ["test9", "test11", "test12", "test25", "test26"]
- # users = ["test9", "test11", ]
- # users = ['test12', 'test25']
- # users = ["test20", "test27"]
- # users = ['test']
- users = [
- "test1",
- "test11",
- "test12",
- "test16",
- "test17",
- "test19",
- "test20",
- "test21",
- "test22",
- "test25",
- "test26",
- "test27",
- "test29",
- "test3",
- "test7",
- "test8",
- "test9",
- ]
- users = ["test"]
- users = ["test12", "test17", "test21", "test22", "test27", ]
- users = ["test27"]
- acc_list = []
- for user in users:
- acc = user_label_accuracy(user)
- acc_list.append([user, acc])
- print(acc_list)
- # get_single_result(161927)
- # import pandas as pd
- # df = pd.read_csv("C:\\Users\\Administrator\\Desktop\\4.csv")
- # _dict = {
- # 51: "公告变更",
- # 52: "招标公告",
- # 101: "中标信息",
- # 102: "招标预告",
- # 103: "招标答疑",
- # 104: "招标文件",
- # 105: "资审结果",
- # 106: "法律法规",
- # 107: "新闻资讯",
- # 108: "拟建项目",
- # 109: "展会推广",
- # 110: "企业名录",
- # 111: "企业资质",
- # 112: "全国工程人员",
- # 113: "业主采购",
- # 114: "采购意向",
- # 115: "拍卖出让",
- # 116: "土地矿产",
- # 117: "产权交易",
- # 118: "废标公告",
- # 119: "候选人公示",
- # 120: "合同公告",
- # }
- # data_list = []
- # for index, row in df.iterrows():
- # if index % 100000 == 0:
- # print("Loop", index)
- # print(_dict[int(row['docchannel'])])
- # print(df.iloc[index, 2])
- # data = row.tolist()
- # data[2] = _dict[int(data[2])]
- # data_list.append(data)
- # df = pd.DataFrame(data_list)
- # df.columns = ['docid', '项目名称', '信息类型', '发布时间', '地区', '业主',
- # '预算金额', '中标供应商', '成交金额', '代理机构']
- # df.to_csv("C:\\Users\\Administrator\\Desktop\\4-1.csv", index=False)
|