SPIRAL: Learning to Search and Aggregate
Idea
Mismatch between how reasoning models are trained and how people actually spend extra inference compute at test time.
GRPO trains the model to produce a strong single reasoning trace, but with a lot of compute people typically try parallel strategies and then combine useful ones and then repeat.
Thus, SPIRAL trains on that second procedure. Objective:
Here is the original problem, are independent search traces, and is the later aggregation trace. scores only the final aggregation answer. Same model is used for both stages.
Key point is that the search traces are rewarded for helping the later aggregator produce a correct answer (not just for being correct).
Three types of inference compute: sequential + parallel + aggregative. It jointly learns all three.
Basically, the issue with final-answer-only GRPO is you can’t directly reward exploration and partial progress → too sparse rewards. Uses set RL.
Set reinforcement learning
Scores a collection of samples jointly, . In SPIRAL, the set score is:
Feed to the aggregator. Measure the expected reward of what the aggregator produces.
Then the simplest policy gradient estimator would give every trace in the set the same advantage:
is the advantage for a set. This is noisy though, so they sample a large pool of traces and construct overlapping subsets.
Their implementation
For each maths problem:
1. Sample eight search traces
Ordinary CoT from the prompt.
2. Construct four sets of four traces
From the eight traces, construct , where each contains four distinct search traces. Overlap lowers noise on whether a certain trace is helpful.
These are picked uniformly at random, without replacement, from the available sets.
3. Sample four aggregation traces for each set
So 16 aggregation traces. The aggregation prompt tells the model to audit each candidate trace, identify useful supported ideas, identify errors and unsupported leaps, verify the useful components, and synthesise a final solution.
4. Reward the final aggregation
16 rewards.
5. Advantage for the aggregation traces
The aggregations get ordinary group-relative advantages with respect to the other aggregations from the same :
The model gets better at aggregating what it gets given.
6. Advantage for each set
For a set , SPIRAL averages the reward of its four aggregations (), then takes the average across all four sets to get group-relative advantages for each :
7. Advantage for each search trace
Suppose search trace appeared in . Then SPIRAL gives the average advantage of the sets containing it:
Thus:
They can apply it recursively too.
Issue: the sets are not ordered, but the attention is. Therefore you would actually have to sample from rather than to avoid bias.
Gradient
Set RL gradient
+ Standard RL gradient
Thoughts
I like this, but I wonder whether it’s really needed anymore with agentic systems that can launch sub-agents during RL, since this seems kind of overlapping. Or maybe we could fuse the two together and do this within the large-scale agentic stuff somehow.