MCPcopy Create free account
hub / github.com/Alibaba-NLP/DeepResearch / main

Function main

WebAgent/ParallelMuse/compressed_reasoning_aggregation.py:252–291  ·  view source on GitHub ↗
()

Source from the content-addressed store, hash-verified

250
251
252async 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
294if __name__ == '__main__':

Calls 3

cluster_by_questionFunction · 0.85
call_convergeFunction · 0.85
read_jsonlFunction · 0.70

Tested by

no test coverage detected