| 49 | |
| 50 | |
| 51 | def package_sqls(sql_path, db_root_path, mode='gpt', data_mode='dev'): |
| 52 | clean_sqls = [] |
| 53 | db_path_list = [] |
| 54 | if mode == 'gpt': |
| 55 | sqls = open(sql_path) |
| 56 | sql_txt = sqls.readlines() |
| 57 | for index, sql_str in enumerate(sql_txt): |
| 58 | if isinstance(sql_str, str): |
| 59 | try: |
| 60 | sql, db_name = sql_str.strip().split('\t----- bird -----\t') |
| 61 | # sql, db_name = sql_str.strip().split('\t') |
| 62 | except: |
| 63 | sql, db_name = " ", "financial" |
| 64 | clean_sqls.append(sql) |
| 65 | db_path_list.append(db_root_path + db_name + '/' + db_name + '.sqlite') |
| 66 | sqls.close() |
| 67 | # print(clean_sqls) |
| 68 | elif mode == 'gt': |
| 69 | sqls = open(sql_path) |
| 70 | sql_txt = sqls.readlines() |
| 71 | # sql_txt = [sql.split('\t')[0] for sql in sql_txt] |
| 72 | for idx, sql_str in enumerate(sql_txt): |
| 73 | sql, db_name = sql_str.strip().split('\t') |
| 74 | clean_sqls.append(sql) |
| 75 | db_path_list.append(db_root_path + db_name + '/' + db_name + '.sqlite') |
| 76 | sqls.close() |
| 77 | |
| 78 | return clean_sqls, db_path_list |
| 79 | |
| 80 | |
| 81 | def run_sqls_parallel(sqls, db_places, num_cpus=1, meta_time_out=30.0): |