MCPcopy Create free account
hub / github.com/MARIO-Math-Reasoning/Super_MARIO / solve

Method solve

mcts_math/solver.py:206–242  ·  view source on GitHub ↗
(self, solvers: List[BaseTree])

Source from the content-addressed store, hash-verified

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 = {}

Callers 1

solver_demo.pyFile · 0.80

Calls 6

generate_preprocessMethod · 0.95
generate_postprocessMethod · 0.95
value_preprocessMethod · 0.95
value_postprocessMethod · 0.95
postprocessMethod · 0.95
outputMethod · 0.95

Tested by

no test coverage detected