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)