- FLAIR: Future Likelihood Advantage for Intermediate Reasoning
- 1. Motivation
- 2. Problem Setup
- 3. Basic FLAIR Score
- 4. Algorithm 1: Basic Candidate Ranking
- 5. Algorithm 2: FLAIR Preference Pair Construction
- 6. Algorithm 3: Block-Level FLAIR for Long Rollouts
- 7. Using GRWO to Generate Better FLAIR Samples
- 8. Training Objectives
- 9. Practical Target Design
- 10. Diagnostics and Evaluation
- 11. Limitations
- 12. Recommended Default Configuration
- 13. Expected Effects
- 14. Minimal Implementation Checklist
- 15. Summary
- 1. Motivation
FLAIR: Future Likelihood Advantage for Intermediate Reasoning
FLAIR is an experimental credit-assignment method for reasoning-model training.
The main idea is simple:
A reasoning chunk is useful if it makes the correct future solution more likely.
Instead of only asking whether a generated chunk "looks correct", FLAIR scores whether the chunk improves the model's likelihood of a known target such as a final answer, key solution step, verifier-approved solution, or reference continuation.
This README describes the basic FLAIR method, its algorithm, how it can be used to create training pairs, and how GRWO can generate better FLAIR samples.
1. Motivation
Reasoning traces are messy.
A model may generate a 512-token reasoning window that contains:
restatement -> weak exploration -> wrong idea -> self-correction -> key insight -> correct continuation
If this whole window eventually leads to the correct answer, standard preference training may treat the whole window as good. That is a credit-assignment problem.
The opposite can also happen:
correct setup -> generic ramble -> algebra mistake -> wrong conclusion
If the final answer is wrong, standard preference training may punish the whole window, including the useful setup.
This is especially bad for math reasoning, where a few tokens can determine the whole trajectory:
Since x > 0, the sign depends on x^2 - a.
or
Use the bipartition invariant.
or
Let d = ab/(a+b), so (a-d)(b-d)=d^2.
FLAIR tries to isolate these useful reasoning moves.
The goal is not just to reward full answers. The goal is to identify which intermediate reasoning chunks make the correct future easier to produce.
2. Problem Setup
We assume we have a prompt, a partial reasoning prefix, one or more candidate chunks, and a target.
Definitions:
P = original problem prompt
R = current reasoning prefix
C = candidate reasoning chunk
T = target future text
M = scoring model
The target T can be one of several forms:
1. Final answer only
2. Short key solution step
3. Reference solution prefix
4. Full reference solution
5. Verifier-approved answer explanation
6. Gold notes containing key facts and final answer
The best target is usually not a huge full solution. A long target can reward surface imitation instead of true reasoning.
A better target often looks like:
keypoint + compact derivation + final answer
Example target for a derivative-sign problem:
Since x > 0, f'(x) = (x^2-a)/x, so the sign depends on x^2-a.
For f(x) >= 0 on [1,infty), all a < 0 work and for a > 0 we need a <= 1.
Answer: a in (-infty,0) union (0,1].
FLAIR then asks:
Does adding candidate chunk C after prefix R increase log probability of target T?
3. Basic FLAIR Score
The basic score compares the target likelihood before and after adding a reasoning chunk.
Plain formula:
base_score = logprob_M(T | P + R)
after_score = logprob_M(T | P + R + C)
delta = after_score - base_score
Interpretation:
delta > 0 -> C made the target more likely
delta = 0 -> C had little effect
delta < 0 -> C made the target less likely
This gives a future-facing score.
A chunk is good not because it is locally fluent, but because it improves the probability of the correct continuation.
Example
Problem:
Given f(x)=1/2 x^2 - a ln x - 1/2, find monotonicity and the range of a
such that f(x)>=0 for x in [1,infty).
Prefix:
We have f'(x)=x-a/x.
Good chunk:
Since x>0, f'(x)=(x^2-a)/x, so the sign depends on x^2-a.
Bad chunk:
Multiplying by x^2 gives x^3-a, so the sign depends on x^3-a.
Both chunks are fluent. Only the first chunk helps the correct future. FLAIR should assign a higher delta to the good chunk.
4. Algorithm 1: Basic Candidate Ranking
This is the simplest version of FLAIR.
Algorithm 1: Basic FLAIR Candidate Ranking
Input:
problem prompt P
reasoning prefix R
candidate chunks C_1 ... C_n
target text T
scoring model M
Step 1:
Compute base_score:
base_score = logprob_M(T | P + R)
Step 2:
For each candidate chunk C_i:
after_score_i = logprob_M(T | P + R + C_i)
delta_i = after_score_i - base_score
Step 3:
Rank candidates by delta_i.
Output:
best_chunk = candidate with highest delta
worst_chunk = candidate with lowest delta
all scored candidates
Minimal pseudocode:
def flair_rank(model, prompt, prefix, chunks, target):
base_score = logprob(model, context=prompt + prefix, target=target)
scored = []
for chunk in chunks:
after_score = logprob(model, context=prompt + prefix + chunk, target=target)
delta = after_score - base_score
scored.append({
"chunk": chunk,
"base_score": base_score,
"after_score": after_score,
"delta": delta,
})
scored.sort(key=lambda x: x["delta"], reverse=True)
return scored
A common output format:
{
"problem": "...",
"prefix": "...",
"chunk": "...",
"base_score": -0.3006,
"after_score": -0.2801,
"delta": 0.0205,
"target_token_count": 589
}
5. Algorithm 2: FLAIR Preference Pair Construction
FLAIR can create preference pairs for DPO, ORPO, KTO-style filtering, or custom preference objectives.
The simplest pair construction is:
chosen = highest-delta chunk
rejected = lowest-delta chunk
But in practice, we should add filters.
Recommended pair rules:
1. chosen_delta must be greater than a positive threshold
2. rejected_delta should be negative or clearly lower than chosen_delta
3. chosen - rejected margin should be large enough
4. chosen chunk should not be pure filler
5. rejected chunk should not contain a correct key step that is only followed by later drift
Plain algorithm:
Algorithm 2: FLAIR Pair Builder
Input:
scored chunks from Algorithm 1
positive threshold tau_pos
margin threshold tau_margin
Step 1:
chosen = chunk with maximum delta
rejected = chunk with minimum delta
Step 2:
Check:
chosen.delta > tau_pos
chosen.delta - rejected.delta > tau_margin
Step 3:
If checks pass:
emit preference pair
Output:
preference pair:
prompt = P + R
chosen = chosen.chunk
rejected = rejected.chunk
Pseudocode:
def build_flair_pair(scored, tau_pos=0.0, tau_margin=0.01):
scored = sorted(scored, key=lambda x: x["delta"], reverse=True)
chosen = scored[0]
rejected = scored[-1]
if chosen["delta"] <= tau_pos:
return None
if chosen["delta"] - rejected["delta"] <= tau_margin:
return None
return {
"chosen": chosen["chunk"],
"rejected": rejected["chunk"],
"chosen_delta": chosen["delta"],
"rejected_delta": rejected["delta"],
"margin": chosen["delta"] - rejected["delta"],
}
This creates a training sample like:
{
"prompt": "Problem + current reasoning prefix",
"chosen": "Since x>0, f'(x)=(x^2-a)/x...",
"rejected": "Multiplying by x^2 gives x^3-a..."
}
This is already stronger than random chosen/rejected generation because the pair is selected by downstream target likelihood.
6. Algorithm 3: Block-Level FLAIR for Long Rollouts
A full 512-token rollout can contain mixed reasoning.
Example:
tokens 0-128: generic restatement
tokens 128-256: wrong approach
tokens 256-384: self-correction
tokens 384-512: correct key insight
If we train on the whole 512 tokens, we may reinforce filler and mistakes.
Block-level FLAIR solves this by scoring parts of the rollout.
Recommended default:
rollout length = 512 tokens
block size = 64 tokens
stride = 64 tokens
training span = 128 to 192 tokens
Block score:
score_before_block = logprob(T | P + R + rollout before block)
score_after_block = logprob(T | P + R + rollout through block)
block_delta = score_after_block - score_before_block
Plain algorithm:
Algorithm 3: Block-Level FLAIR
Input:
prompt P
prefix R
rollout W
target T
block size B
Step 1:
Split rollout W into blocks:
W_1, W_2, ..., W_k
Step 2:
For each block W_i:
context_before = P + R + W_1 + ... + W_{i-1}
context_after = P + R + W_1 + ... + W_i
before_score = logprob(T | context_before)
after_score = logprob(T | context_after)
block_delta = after_score - before_score
Step 3:
Select:
best_block = block with highest block_delta
worst_block = block with lowest block_delta
Output:
block-level credit map
Pseudocode:
def block_flair(model, prompt, prefix, rollout_blocks, target):
results = []
previous_context = prompt + prefix
before_score = logprob(model, previous_context, target)
for idx, block in enumerate(rollout_blocks):
after_context = previous_context + block
after_score = logprob(model, after_context, target)
delta = after_score - before_score
results.append({
"block_id": idx,
"block": block,
"before_score": before_score,
"after_score": after_score,
"delta": delta,
})
previous_context = after_context
before_score = after_score
return results
This lets us detect a late breakthrough.
If the model yaps for 300 tokens but finds the real solution at token 380, block-level FLAIR can still select the breakthrough span.
7. Using GRWO to Generate Better FLAIR Samples
GRWO and FLAIR solve different parts of the problem.
GRWO generates local candidate continuations using a guide.
FLAIR scores which continuation actually improves the future target.
A good pipeline is:
GRWO = generate candidate paths
FLAIR = score which path helps
DPO/RL = train on selected path
GRWO-assisted FLAIR pipeline
Algorithm 4: GRWO-Assisted FLAIR Data Generation
Input:
problem P
private solution guide G
base model M
scoring model S
Step 1:
Generate a reasoning prefix R from P.
Step 2:
Generate candidate windows:
C_unguided = model continuation without guide
C_guided = model continuation using private guide
C_repaired = continuation after correcting a bad path
C_negative = intentionally weak or unguided continuation
Step 3:
Score every candidate with FLAIR:
delta_i = logprob(T | P + R + C_i) - logprob(T | P + R)
Step 4:
Select:
chosen = high-delta candidate/span
rejected = low-delta candidate/span
Step 5:
Emit training pair.
Pseudocode:
def grwo_flair_sample(generator, scorer, problem, guide, target):
prefix = generator.generate_prefix(problem)
candidates = []
candidates.append(generator.generate(problem, prefix, guide=None))
candidates.append(generator.generate(problem, prefix, guide=guide))
candidates.append(generator.generate_repair(problem, prefix, guide=guide))
candidates.append(generator.generate_negative(problem, prefix))
scored = flair_rank(
model=scorer,
prompt=problem,
prefix=prefix,
chunks=candidates,
target=target,
)
pair = build_flair_pair(scored)
return pair
Why this is better than plain GRWO
Plain GRWO can be too local:
prefix -> guided continuation -> DPO update
But the guided continuation may only look better locally. It may not actually improve the final solution.
FLAIR adds a future check:
Does this guided continuation increase probability of the correct target?
This helps remove bad or superficial pairs.
Recommended generation settings
For a small reasoning model, a good starting setup is:
candidate_count = 4 to 8
prefix_tokens = 32 to 96
candidate_tokens = 256 to 512
block_size = 64
training_span_tokens = 128 to 192
temperature = 0.7 for exploration, 0.0 for deterministic validation
For stable pair generation:
generate long enough to reveal trajectory
train only the span with useful credit
8. Training Objectives
FLAIR is a scoring method. It can be plugged into several training objectives.
8.1 DPO-style training
Use FLAIR to create chosen/rejected pairs.
chosen = high future-likelihood chunk
rejected = low future-likelihood chunk
The DPO objective then trains the model to prefer chosen over rejected.
Plain objective intuition:
increase log probability of chosen over rejected
relative to a reference model
Recommended for early experiments because it is simple and stable.
8.2 Weighted DPO
If block-level FLAIR is available, weight tokens near high-impact blocks more strongly.
Example:
weight = 1.0 + clipped_positive_block_delta
Then the loss focuses more on useful parts of the reasoning window.
8.3 GRPO/PPO-style training
FLAIR can also define a rollout reward.
For each rollout:
reward = final_answer_reward + alpha * flair_delta
Example:
final_answer_reward = 1 if answer correct else 0
flair_delta = logprob target after rollout - logprob target before rollout
Then standardize rewards inside a group of sampled completions:
advantage_i = (reward_i - mean_reward) / (std_reward + eps)
This gives a group-relative signal similar to GRPO.
8.4 Span-level policy update
Instead of training the whole rollout, train only a span selected by FLAIR.
selected_span = span around highest positive block_delta
This is useful because DPO and GRPO often give blunt credit to too many tokens.
9. Practical Target Design
Target design is the most important part of FLAIR.
Bad target:
A very long reference solution with lots of style-specific wording.
This can make the scorer reward chunks that imitate wording instead of chunks that solve the problem.
Better target:
A compact target containing:
1. the key mathematical move
2. the main derivation
3. the final answer
For math problems, recommended target formats:
Format A: final answer only
Answer: 120
Pros:
cheap, direct
Cons:
weak for long reasoning, may miss useful intermediate steps
Format B: keypoint target
Key step: treat EE and SS as blocks.
Answer: 5! = 120.
Pros:
strong signal for reasoning move
Cons:
requires keypoint extraction
Format C: compact verifier notes
Use Lagrange multipliers or generalized eigenvalue method.
The maximum occurs at the largest generalized eigenvalue times 22.
Final answer: ...
Pros:
more robust than final answer only
Cons:
more expensive to prepare
Format D: multi-target scoring
Score several targets and average them:
target_1 = key step
target_2 = compact derivation
target_3 = final answer
This reduces wording overfit.
Recommended score:
flair_score = average(delta over targets)
10. Diagnostics and Evaluation
FLAIR should not be trusted blindly. Track whether it correlates with actual correctness.
Useful metrics:
1. Average delta for known good chunks
2. Average delta for known bad chunks
3. Good-vs-bad separation
4. Spearman correlation between label and delta
5. Best good chunk > best bad chunk rate
6. Every good chunk > every bad chunk rate
7. Final answer accuracy after training
8. Keypoint hit rate
9. Wrong-pivot frequency
10. Repetition rate
Example diagnostic report:
Average delta_logprob, good chunks: 1.0155
Average delta_logprob, bad chunks: -0.4783
Average delta_margin, good chunks: 1.3726
Average delta_margin, bad chunks: -0.8954
Best good > best bad by margin: 80.0%
Spearman(label, delta_margin): 0.7311
This kind of result suggests that the scoring signal is meaningful.
What good FLAIR behavior looks like
For good chunks:
delta should be positive
target logprob should improve
correct keypoint should become easier to generate
For bad chunks:
delta should be negative
wrong continuation should become more likely
correct target should become less likely
Red flags
FLAIR may be noisy if:
good and bad chunks have similar delta
generic phrases score too highly
long reference targets dominate the score
answer-only target gives unstable signal
negative chunks contain correct setup
positive chunks contain hidden algebra mistakes
11. Limitations
FLAIR is not a complete solution by itself.
11.1 It can reward reference similarity
If the target is too close to a reference solution, chunks with similar wording can score higher even if they are not mathematically causal.
Mitigation:
use compact keypoint targets
use multiple target phrasings
use final answer verification
11.2 It can miss non-verbal insight
Some chunks improve reasoning structure without directly overlapping target text.
Mitigation:
score multiple future targets
include keypoint and answer targets
evaluate with actual generation after the chunk
11.3 It is compute-heavy
Scoring many chunks against targets requires additional forward passes.
Mitigation:
score blocks instead of tokens
use 64-token block checkpoints
cache prefix KV states when possible
score only top candidate windows
11.4 It can train bad spans if chunk selection is too broad
A high-scoring 512-token rollout may contain bad sections.
Mitigation:
use block-level FLAIR
train 128-192 token spans
filter obvious invalid math
11.5 It does not replace final verification
A chunk can improve target likelihood but still lead to a wrong final answer.
Mitigation:
combine FLAIR with final answer reward
use symbolic verifiers where possible
track benchmark accuracy separately
12. Recommended Default Configuration
A practical starting configuration:
FLAIR_CONFIG = {
"candidate_count": 8,
"prefix_tokens": 64,
"rollout_tokens": 512,
"block_size": 64,
"training_span_tokens": 192,
"target_mode": "keypoint_plus_answer",
"score_mode": "avg_logprob",
"pair_margin_threshold": 0.01,
"positive_delta_threshold": 0.0,
}
For cheaper experiments:
FLAIR_CONFIG_CHEAP = {
"candidate_count": 4,
"prefix_tokens": 48,
"rollout_tokens": 256,
"block_size": 64,
"training_span_tokens": 128,
}
For stronger experiments:
FLAIR_CONFIG_STRONG = {
"candidate_count": 16,
"prefix_tokens": 64,
"rollout_tokens": 512,
"block_size": 32,
"training_span_tokens": 192,
"multi_target_scoring": True,
}
Recommended first training setup:
generate 512 tokens
score by 64-token blocks
select 128-192 token high-impact span
pair with low-impact or harmful span
train with DPO or weighted DPO
13. Expected Effects
If FLAIR works, the model should improve in reasoning dynamics before accuracy fully catches up.
Expected changes:
1. better key-step selection
2. fewer wrong early branches
3. stronger self-correction
4. less useless restatement
5. better handling of invariants
6. better derivative sign analysis
7. better parameter range reasoning
8. better ability to recover from weak starts
The most important expected change:
The model becomes better at finding the decisive next reasoning move.
Not every benchmark will improve immediately. First, the trajectories should become cleaner. Then final accuracy should improve.
14. Minimal Implementation Checklist
To implement basic FLAIR:
[ ] Load scoring model
[ ] Prepare prompt P
[ ] Prepare prefix R
[ ] Generate candidate chunks C_i
[ ] Prepare target T
[ ] Compute base target logprob
[ ] Compute after target logprob for each candidate
[ ] Compute delta for each candidate
[ ] Rank candidates by delta
[ ] Build chosen/rejected pairs
[ ] Train or filter dataset
[ ] Evaluate correlation with correctness
Common mistakes:
[ ] using full reference only as target
[ ] training whole 512-token rollout without block scoring
[ ] ignoring negative chunks that contain useful setup
[ ] using too many noisy samples without margin filtering
[ ] assuming positive delta always means mathematically correct
15. Summary
FLAIR is a future-likelihood credit-assignment method for intermediate reasoning.
Core principle:
Reward reasoning chunks that make the correct future solution more likely.
Basic FLAIR:
delta = logprob(target | prefix + chunk) - logprob(target | prefix)
Use cases:
1. rank candidate reasoning chunks
2. filter reasoning data
3. build chosen/rejected pairs
4. select training spans
5. improve GRWO sample quality
6. provide reward signals for RL-style training
Best practical rule:
Generate long enough to see the path.
Score the future impact.
Train only the span that caused the improvement.
FLAIR is not a magic verifier. It is a way to make credit assignment less blunt.
The method is strongest when combined with:
- compact keypoint targets
- final answer verification
- block-level scoring
- GRWO candidate generation
- margin filtering
- careful diagnostics