(self, questions: List[str])
| 81 | return solver |
| 82 | |
| 83 | def batch_generate(self, questions: List[str]): |
| 84 | |
| 85 | solvers = [REACT(config=self.config, question=question) for question in questions] |
| 86 | |
| 87 | for step in tqdm(range(self.config.max_depth), desc="Step Processing"): |
| 88 | prompts = [] |
| 89 | epoch_solvers = [] |
| 90 | next_solvers = [] |
| 91 | |
| 92 | for solver in solvers: |
| 93 | if solver.should_generate_next(): |
| 94 | prompts.append(solver.create_prompt()) |
| 95 | epoch_solvers.append(solver) |
| 96 | else: |
| 97 | next_solvers.append(solver) |
| 98 | |
| 99 | next_solver_span = len(next_solvers) |
| 100 | if len(epoch_solvers) < 1: |
| 101 | break |
| 102 | |
| 103 | # vllm run |
| 104 | outputs = self.llm(prompts) |
| 105 | # post-process outputs |
| 106 | with ProcessPool(max_workers=min(len(epoch_solvers), os.cpu_count())) as pool: |
| 107 | future = pool.map(self.__class__.processor, epoch_solvers, outputs, timeout=TIMEOUT_SECONDS) |
| 108 | iterator = future.result() |
| 109 | |
| 110 | if len(epoch_solvers) > 100: |
| 111 | progress_bar = tqdm(total=len(epoch_solvers), desc="Execute") |
| 112 | else: |
| 113 | progress_bar = None |
| 114 | |
| 115 | while True: |
| 116 | try: |
| 117 | result = next(iterator) |
| 118 | next_solvers.append(result) |
| 119 | except StopIteration: |
| 120 | break |
| 121 | except TimeoutError as error: |
| 122 | next_solvers.append(None) |
| 123 | print(error) |
| 124 | except Exception as error: |
| 125 | print(error) |
| 126 | next_solvers.append(None) |
| 127 | if progress_bar is not None: |
| 128 | progress_bar.update(1) |
| 129 | |
| 130 | if progress_bar is not None: |
| 131 | progress_bar.close() |
| 132 | |
| 133 | # update solvers |
| 134 | assert len(epoch_solvers) == len(next_solvers[next_solver_span:]), f"Data is not matched, {len(epoch_solvers)} vs {len(next_solvers[next_solver_span:])}." |
| 135 | for idx, (ori_solver, new_solver) in enumerate(zip(epoch_solvers, next_solvers[next_solver_span:])): |
| 136 | if new_solver is None: |
| 137 | next_solvers[next_solver_span + idx] = ori_solver |
| 138 | solvers = next_solvers |
| 139 | |
| 140 | jsonlines = {} |
no test coverage detected