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
Downloads last month

-

Downloads are not tracked for this model. How to track
Inference Providers NEW
This model isn't deployed by any Inference Provider. 🙋 Ask for provider support