SPIRAL: Learning to Search and Aggregate

Paper authors: Jubayer Ibn Hamid, Ifdita Hasan Orney, Michael Y. Li, Omar Shaikh, Yoonho Lee, Dorsa Sadigh, Chelsea Finn, and Noah Goodman.

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:

J(θ)=𝔼y1:n~πθ(·∣x)[𝔼y*~πθ(·∣x,y1:n)[r(x,y*)]].

Here x is the original problem, y1,…,yn are independent search traces, and y* is the later aggregation trace. r(x,y*) 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, f(x,y1:n). In SPIRAL, the set score is:

fSPIRAL(x,y1:n)=𝔼y*~πθ(·∣x,y1:n)[r(x,y*)].

Feed y1:n 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:

∇θJsearch≈A♯(x,y1:n)∑i=1n∇θlogπθ(yi∣x).

A♯ 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

y1,…,y8~i.i.d.πθ(·∣x).

Ordinary CoT from the prompt.

2. Construct four sets of four traces

From the eight traces, construct G1,G2,G3,G4, where each Gi 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 (84) available sets.

3. Sample four aggregation traces for each set

y1Gi,…,y4Gi~πθ(·∣x,Gi).

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 Gi:

A(x,yjGi)=r(x,yjGi)−r¯Gi,
r¯Gi=14∑j=14r(x,yjGi).

The model gets better at aggregating what it gets given.

6. Advantage for each set

For a set Gi, SPIRAL averages the reward of its four aggregations (r¯Gi), then takes the average across all four sets to get group-relative advantages for each Gi:

A♯(x,Gi)=r¯Gi−14∑k=14r¯Gk.

7. Advantage for each search trace

Suppose search trace y1 appeared in G1,G2,G3. Then SPIRAL gives y1 the average advantage of the sets containing it:

Amarg♯(x,y1)=A♯(x,G1)+A♯(x,G2)+A♯(x,G3)3.

Thus:

∇θJsearch≈∑i=18Amarg♯(x,yi)∇θlogπθ(yi∣x).

They can apply it recursively too.

Issue: the sets Gi are not ordered, but the attention is. Therefore you would actually have to sample from 8P4 rather than (84) to avoid bias.

Gradient

∇θJ(θ)=∇θ𝔼y1:n~πθ(·∣x)[𝔼y*~πθ(·∣x,y1:n)[r(x,y*)]].
∇θJ(θ)=𝔼y1:n~πθ(·∣x)𝔼y*~πθ(·∣x,y1:n)[r(x,y*)∇θlog(πθ(y1:n∣x)πθ(y*∣x,y1:n))].
=𝔼y1:n~πθ(·∣x)𝔼y*~πθ(·∣x,y1:n)[r(x,y*)(∇θlogπθ(y1:n∣x)+∇θlogπθ(y*∣x,y1:n))].

Set RL gradient

=𝔼y1:n~πθ(·∣x)[∇θlogπθ(y1:n∣x)𝔼y*~πθ(·∣x,y1:n)[r(x,y*)]]

+ Standard RL gradient

+𝔼y1:n~πθ(·∣x)[𝔼y*~πθ(·∣x,y1:n)[∇θlogπθ(y*∣x,y1:n)r(x,y*)]].

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.