| 474 | return input_dict |
| 475 | |
| 476 | def save_sql_txt(self, save_file_name, suffix_name, sql, dataset, db_id): |
| 477 | filename, ext = os.path.splitext(save_file_name) |
| 478 | filename = filename + suffix_name |
| 479 | new_path = filename + ext |
| 480 | new_path = os.path.join("outputs", dataset, new_path) |
| 481 | |
| 482 | with open(new_path, "a+") as f: |
| 483 | if dataset == "spider": |
| 484 | f.write(sql.replace("\n", " ")) |
| 485 | f.write("\n") |
| 486 | elif dataset == "bird": |
| 487 | f.write((sql.replace("\n", " ") + '\t----- bird -----\t' + 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) |