MCPcopy Create free account
hub / github.com/FlyingFeather/DEA-SQL / save_file

Method save_file

gen_sql.py:490–537  ·  view source on GitHub ↗
(self, save_file_name, dataset, final_sql, final_correct_sql, final_des_sql, mode, item, db_id)

Source from the content-addressed store, hash-verified

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"):

Callers 1

mainMethod · 0.95

Calls 3

save_sql_txtMethod · 0.95
evaluateFunction · 0.90

Tested by

no test coverage detected