()
| 250 | |
| 251 | |
| 252 | async def main(): |
| 253 | data_path = "[YOUR-ROLLOUT-FILE-PATH-HERE]" # TODO |
| 254 | |
| 255 | report_sem = asyncio.Semaphore(64) |
| 256 | merge_sem = asyncio.Semaphore(32) |
| 257 | sem = { |
| 258 | 'report': report_sem, |
| 259 | 'merge': merge_sem |
| 260 | } |
| 261 | mode = 'converge_info' |
| 262 | |
| 263 | dataset = read_jsonl(data_path) |
| 264 | |
| 265 | tasks = [] |
| 266 | clustered_dataset = cluster_by_question(dataset) |
| 267 | for cluster in clustered_dataset: |
| 268 | filtered_cluster = [] |
| 269 | for traj in cluster: |
| 270 | if 'prediction' in traj.keys() and traj['prediction'] != '[No Prediction]': |
| 271 | filtered_cluster.append(traj) |
| 272 | |
| 273 | tasks.append(call_converge(sem, filtered_cluster, data_path)) |
| 274 | |
| 275 | results = [] |
| 276 | |
| 277 | with open(f"{data_path.replace('.jsonl', f'_{mode}.jsonl')}", "a") as f: |
| 278 | for future in tqdm(asyncio.as_completed(tasks), total=len(tasks), desc=f"Converging ..."): |
| 279 | try: |
| 280 | result = await future |
| 281 | results.append(result) |
| 282 | f.write(json.dumps(result, ensure_ascii=False) + "\n") |
| 283 | f.flush() |
| 284 | os.fsync(f.fileno()) |
| 285 | except Exception as e: |
| 286 | exception_type = type(e).__name__ |
| 287 | exception_message = str(e) |
| 288 | traceback_info = ''.join(traceback.format_tb(e.__traceback__)) |
| 289 | error_message = f'{exception_type}: {exception_message}\n' \ |
| 290 | f'Traceback:\n{traceback_info}' |
| 291 | print(f"[ERROR]: {error_message}") |
| 292 | |
| 293 | |
| 294 | if __name__ == '__main__': |
no test coverage detected