| 77 | |
| 78 | |
| 79 | def package_sqls(sql_path, db_root_path, mode='gpt', data_mode='dev'): |
| 80 | clean_sqls = [] |
| 81 | db_path_list = [] |
| 82 | if mode == 'gpt': |
| 83 | sql_data = json.load(open(sql_path, 'r')) |
| 84 | for idx, sql_str in sql_data.items(): |
| 85 | if isinstance(sql_str, str): |
| 86 | sql, db_name = sql_str.split('\t----- bird -----\t') |
| 87 | else: |
| 88 | sql, db_name = " ", "financial" |
| 89 | clean_sqls.append(sql) |
| 90 | db_path_list.append(db_root_path + db_name + '/' + db_name + '.sqlite') |
| 91 | |
| 92 | elif mode == 'gt': |
| 93 | sqls = open(sql_path) |
| 94 | sql_txt = sqls.readlines() |
| 95 | for idx, sql_str in enumerate(sql_txt): |
| 96 | sql, db_name = sql_str.strip().split('\t') |
| 97 | clean_sqls.append(sql) |
| 98 | db_path_list.append(db_root_path + db_name + '/' + db_name + '.sqlite') |
| 99 | |
| 100 | return clean_sqls, db_path_list |
| 101 | |
| 102 | |
| 103 | def run_sqls_parallel(sqls, db_places, num_cpus=1, iterate_num=100, meta_time_out=30.0): |