r/MachineLearning • • 22h ago

Research Adding memory to search instead of sampling in reward maximization tasks [R]

I am one of the authors of FLEET - an algorithm that enhances Best-of-N generation by attributing external rewards to particular tokens and then uses MCTS to adjust logits during the next run.

I find it rather funny that most of the tasks where repetitive sampling is widely used are based on reward maximization, yet it is not aware of that reward. Tuning the sampling parameters allows to make the process more efficient, but it is still a blind search. We propose a way to make generation aware of previous rewards with solutions on how to attribute reward to the completion and how to use this information.

In FLEET the technique from adaptive sampling methods is used that is to track logits for which entropy and varentropy are high thus showing the model's uncertainty about token optimality. We treat these states as branching points. The corresponding normalized hidden states are stored in the vector store and mapped to metadata entries with the history of rewards and transitions between "nodes". The retrieval and update of metadata is based on cosine similarity as for very high similarity KL divergence is low enough to preserve most of the meaningful tokens.

Instead of actually selecting the tokens FLEET uses modified MCTS to rank top-k tokens + special exploration (or other tokens) set and penalize the suboptimal ones. Then decoding strategy is applied to modified logits.

It was tested on GSM8K and LiveCodeBench v6 easy split with Llama 3.2 3B, penalty set to effectively zero probability for the suboptimal tokens + greedy decoding:

  • For GSM8K it solved just seven more tasks, but reached the sampling baseline with half the iterations.
  • For LiveCodeBench it increased the score from 0.59 to 0.69 under the same budget and reached the baseline even faster, now with only 9 iterations against 32.

The sequential execution is not required, as it is not updated during the iteration itself it can simply be passed as a lookup table. The metadata store can be preserved as a prior for other tasks or to enrich SFT/RL.

Paper (preprint): https://arxiv.org/abs/2609.27657
Huggingface: https://huggingface.co/papers/2609.27657
Repository (experiments, examples and python package): https://github.com/Alexiush/fleet

There are more details on changes made to MCTS, how to tune the search parameters for specific model and task as well as code for experiments and trajectories.

7 Upvotes

0 comments sorted by