(self, save_file_name, dataset, final_sql, final_correct_sql, final_des_sql, mode, item, db_id)
| 488 | f.write("\n") |
| 489 | |
| 490 | def save_file(self, save_file_name, dataset, final_sql, final_correct_sql, final_des_sql, mode, item, db_id): |
| 491 | filename, ext = os.path.splitext(save_file_name) |
| 492 | filename = filename + '_clean' |
| 493 | new_path = filename + ext |
| 494 | new_path = os.path.join("outputs", dataset, new_path) |
| 495 | |
| 496 | filename_2 = filename + '_correct' |
| 497 | new_path_2 = filename_2 + ext |
| 498 | new_path_2 = os.path.join("outputs", dataset, new_path_2) |
| 499 | with open(new_path, "a+") as f: |
| 500 | if dataset == "spider": |
| 501 | f.write(final_sql.replace("\n", " ")) |
| 502 | f.write("\n") |
| 503 | elif dataset == "bird": |
| 504 | f.write((final_sql.replace("\n", " ") + '\t----- bird -----\t' + db_id)) |
| 505 | f.write("\n") |
| 506 | |
| 507 | with open(new_path_2, "a+", encoding='utf-8') as f: |
| 508 | if dataset == "spider": |
| 509 | f.write(final_correct_sql.replace("\n", " ")) |
| 510 | f.write("\n") |
| 511 | |
| 512 | suffix_name = "_des" |
| 513 | self.save_sql_txt(save_file_name, suffix_name, final_des_sql, dataset, db_id) |
| 514 | |
| 515 | ############# single evaluate ################# |
| 516 | # only evaluting exact match needs this argument |
| 517 | if mode == "debug": |
| 518 | from single_eval import build_foreign_key_map_from_json, evaluate |
| 519 | kmaps = None |
| 520 | if args.etype in ['all', 'match']: |
| 521 | assert args.table is not None, 'table argument must be non-None if exact set match is evaluated' |
| 522 | kmaps = build_foreign_key_map_from_json(args.table) |
| 523 | |
| 524 | gold_sql = item['query'] + '\t' + item['db_id'] |
| 525 | pred_sql = final_sql.replace("\n", " ") |
| 526 | final_correct_sql = final_correct_sql.replace("\n", " ") |
| 527 | print(f"gold_sql: {item['query']}") |
| 528 | print(f"pred_sql: {pred_sql}") |
| 529 | evaluate(gold_sql, pred_sql, args.db, args.etype, kmaps, args.plug_value, args.keep_distinct, |
| 530 | args.progress_bar_for_each_datapoint) |
| 531 | print(f"final_correct_sql: {final_correct_sql}") |
| 532 | evaluate(gold_sql, final_correct_sql, args.db, args.etype, kmaps, args.plug_value, args.keep_distinct, |
| 533 | args.progress_bar_for_each_datapoint) |
| 534 | |
| 535 | print(f"final_des_sql: {final_des_sql}") |
| 536 | evaluate(gold_sql, final_des_sql, args.db, args.etype, kmaps, args.plug_value, args.keep_distinct, |
| 537 | args.progress_bar_for_each_datapoint) |
| 538 | |
| 539 | def main(self, root_dir, dataset, file, save_file_name, mode="dev", lang_mode="cn", prompt_mode='openai', |
| 540 | sample="True", data_fold="1", test_id=46, insert_value=0, step_name="all"): |
no test coverage detected