(self, solvers: List[BaseTree])
| 204 | return invalid_solvers |
| 205 | |
| 206 | def solve(self, solvers: List[BaseTree]): |
| 207 | |
| 208 | for step in tqdm(range(self.max_solver_steps), desc="Step Processing"): |
| 209 | |
| 210 | prompts, prompts_span, valid_solvers, invalid_solvers = self.generate_preprocess(solvers) |
| 211 | |
| 212 | if len(valid_solvers) < 1: |
| 213 | break |
| 214 | |
| 215 | # llm run for step generation |
| 216 | if step == 0: |
| 217 | n = self.config.n_generate_sample * self.config.step_beam_width |
| 218 | else: |
| 219 | n = self.config.n_generate_sample |
| 220 | self.generate_sampling_params.n = n |
| 221 | self.generate_sampling_params.best_of = n |
| 222 | |
| 223 | outputs = self.llm(prompts, self.generate_sampling_params) |
| 224 | # post-process outputs |
| 225 | reconstructed_outputs = [outputs[bos_idx : eos_idx] for bos_idx, eos_idx in zip(prompts_span, prompts_span[1:])] |
| 226 | |
| 227 | # process output and run python interpreter |
| 228 | valid_solvers = self.generate_postprocess(reconstructed_outputs, valid_solvers) |
| 229 | |
| 230 | # llm run for step evaluation |
| 231 | prompts, prompts_span = self.value_preprocess(valid_solvers) |
| 232 | if self.need_value_func: |
| 233 | outputs = self.llm(prompts, self.value_sampling_params) |
| 234 | reconstructed_outputs = [outputs[bos_idx : eos_idx] for bos_idx, eos_idx in zip(prompts_span, prompts_span[1:])] |
| 235 | else: |
| 236 | reconstructed_outputs = [None] * (len(prompts_span) - 1) |
| 237 | |
| 238 | valid_solvers = self.value_postprocess(reconstructed_outputs, valid_solvers) |
| 239 | |
| 240 | solvers = self.postprocess(valid_solvers, invalid_solvers) |
| 241 | |
| 242 | return self.output(solvers) |
| 243 | |
| 244 | def output(self, solvers: List[BaseTree]): |
| 245 | jsonlines = {} |
no test coverage detected