| name | stepwiser-generative-judges |
| title | StepWiser Stepwise Generative Judges for Wiser Reasoning |
| version | 0.0.2 |
| engine | skillxiv-v0.0.2-claude-opus-4.6 |
| license | MIT |
| url | https://arxiv.org/abs/2508.19229 |
| keywords | ["process-reward-model","generative-judge","reasoning-verification","meta-reasoning","rl-training"] |
| description | Train stepwise judges as generative models that perform meta-reasoning about intermediate steps, combining explainability with improved accuracy over static process reward models |
StepWiser: Stepwise Generative Judges for Wiser Reasoning
Core Concept
StepWiser reframes process reward modeling as a reasoning task itself. Instead of classifying intermediate steps as correct/incorrect, a generative judge model performs meta-reasoning by generating explanations before predicting verdicts. Trained via reinforcement learning on relative outcomes, StepWiser achieves better step-level accuracy than existing methods while being interpretable and generalizable to new problem distributions.
Architecture Overview
- Generative Judge Model: Produces reasoning before verdict
- Meta-Reasoning: Model explains why steps are good/bad
- RL Training: Learn from relative outcome comparisons
- Bidirectional Feedback: Training-time and inference-time improvements
- Process Supervision: Fine-grained intermediate step evaluation
Implementation Steps
Stage 1: Generative Judge Architecture
Design a model that generates step explanations before verdicts.
import torch
from torch import nn
from typing import Dict, List, Tuple
class GenerativeStepwiseJudge(nn.Module):
"""Judge that generates reasoning before verdicting on steps"""
def __init__(
self,
model_dim: int = 4096,
vocab_size: int = 32000,
max_explanation_len: int = 128
):
super().__init__()
self.model_dim = model_dim
self.vocab_size = vocab_size
self.max_explanation_len = max_explanation_len
self.step_encoder = nn.Sequential(
nn.Linear(256, model_dim),
nn.ReLU(),
nn.Linear(model_dim, model_dim)
)
self.context_encoder = nn.Sequential(
nn.Linear(512, model_dim),
nn.ReLU(),
nn.Linear(model_dim, model_dim)
)
self.explanation_head = nn.LSTM(
input_size=model_dim,
hidden_size=model_dim,
num_layers=2,
batch_first=True
)
self.explanation_decoder = nn.Linear(model_dim, vocab_size)
.verdict_head = nn.Sequential(
nn.Linear( * model_dim, model_dim),
nn.ReLU(),
nn.Linear(model_dim, )
)
() -> :
step_embed = .step_encoder(.embed_text(current_step))
context_embed = .context_encoder(.embed_context(problem, previous_steps))
combined = torch.cat([context_embed.unsqueeze(), step_embed.unsqueeze()], dim=-)
explanation_hidden, _ = .explanation_head(combined)
explanation_logits = .explanation_decoder(explanation_hidden)
explanation_tokens = explanation_logits.argmax(dim=-)
verdict_input = torch.cat([step_embed, explanation_hidden.squeeze()], dim=-)
verdict_logits = .verdict_head(verdict_input)
verdict_probs = torch.nn.functional.softmax(verdict_logits, dim=-)
confidence = - torch.nn.functional.entropy(verdict_probs)
{
: explanation_tokens,
: explanation_logits,
: verdict_logits,
: verdict_probs,
: confidence
}
() -> torch.Tensor:
torch.randn()
() -> torch.Tensor:
torch.randn()
Stage 2: RL Training via Comparative Outcomes
Train judge using RL based on which verdict predictions are correct.
class GenerativeJudgeRLTrainer:
"""Train judge using reinforcement learning"""
def __init__(
self,
judge: GenerativeStepwiseJudge,
policy_model,
lr: float = 1e-5
):
self.judge = judge
self.policy_model = policy_model
self.optimizer = torch.optim.Adam(judge.parameters(), lr=lr)
def generate_step_pair(
self,
problem: str,
previous_steps: List[str]
) -> Tuple[str, str, bool]:
"""
Generate two different step candidates and evaluate which is better.
Returns:
step_a, step_b, label (True if A better)
"""
step_a = self.policy_model.generate_step(problem, previous_steps)
step_b = self.policy_model.generate_step(problem, previous_steps)
outcome_a = self.evaluate_trajectory(problem, previous_steps + [step_a])
outcome_b = self.evaluate_trajectory(problem, previous_steps + [step_b])
label = outcome_a > outcome_b
return step_a, step_b, label
def evaluate_trajectory(self, problem: str, steps: List[str]) -> float:
() -> :
losses = {: , : }
example batch:
problem = example[]
steps = example[]
correct_verdicts = example[]
step_idx, step (steps):
output = .judge(
problem,
steps[:step_idx],
step
)
is_correct = correct_verdicts[step_idx]
gt_verdict = torch.tensor([is_correct], dtype=torch.long)
verdict_loss = torch.nn.functional.cross_entropy(
output[],
gt_verdict
)
losses[] += verdict_loss
explanation_loss =
losses[] += explanation_loss
num_steps = ((ex[]) ex batch)
key losses:
losses[key] /= (num_steps, )
total_loss = (losses.values())
.optimizer.zero_grad()
total_loss.backward()
.optimizer.step()
losses
Stage 3: Inference-Time Step Validation
Use the judge to improve reasoning at inference time.
class ReasoningWithStepValidation:
"""Generate reasoning with stepwise validation"""
def __init__(self, policy_model, judge: GenerativeStepwiseJudge):
self.policy = policy_model
self.judge = judge
def generate_with_validation(
self,
problem: str,
max_steps: int = 20,
temperature: float = 0.7,
validation_threshold: float = 0.5
) -> Dict:
"""
Generate solution step-by-step with validation.
Steps with low judge confidence are regenerated or pruned.
"""
steps = []
confidences = []
all_verdicts = []
for step_idx in range(max_steps):
step = self.policy.generate_step(
problem,
steps,
temperature=temperature
)
verdict_output = self.judge(problem, steps, step)
verdict_prob = verdict_output["verdict_probs"][1]
confidence = verdict_output["confidence"]
if confidence > validation_threshold:
steps.append(step)
confidences.append(confidence.item())
all_verdicts.append(verdict_prob.item())
else:
retry ():
step = .policy.generate_step(
problem,
steps,
temperature=temperature +
)
verdict_output = .judge(problem, steps, step)
confidence = verdict_output[]
confidence > validation_threshold:
steps.append(step)
confidences.append(confidence.item())
all_verdicts.append(verdict_output[][].item())
:
steps.append(step)
confidences.append(confidence.item())
all_verdicts.append()
.policy.is_solution_complete(problem, steps):
{
: steps,
: confidences,
: all_verdicts,
: (confidences) / (confidences) confidences
}
Stage 4: Process Reward Training Integration
Use StepWiser judges to improve policy models.
class ProcessRewardImprovement:
"""Use judge verdicts to improve policy"""
def __init__(self, policy_model, judge: GenerativeStepwiseJudge):
self.policy = policy_model
self.judge = judge
self.policy_optimizer = torch.optim.Adam(policy_model.parameters(), lr=1e-5)
def improve_policy_with_process_rewards(
self,
rollouts: List[Dict]
) -> float:
"""
Train policy to generate steps that pass judge validation.
Reward = judge confidence in step correctness.
"""
total_loss = 0
for rollout in rollouts:
problem = rollout["problem"]
steps = rollout["steps"]
ground_truth = rollout["answer"]
for step_idx, step in enumerate(steps):
log_prob = self.policy.get_log_prob(step, problem, steps[:step_idx])
verdict_output = self.judge(problem, steps[:step_idx], step)
judge_confidence = verdict_output["confidence"]
path_quality = self.evaluate_path_quality(problem, steps[:step_idx+], ground_truth)
reward = * judge_confidence + * path_quality
loss = -(log_prob * reward)
total_loss += loss
total_loss /= ((r[]) r rollouts)
.policy_optimizer.zero_grad()
total_loss.backward()
.policy_optimizer.step()
total_loss.item()
() -> :
Stage 5: Evaluation
Measure judge accuracy and policy improvement.
class StepwiserEvaluator:
"""Evaluate judge quality and reasoning improvements"""
def __init__(self, judge: GenerativeStepwiseJudge, policy_model):
self.judge = judge
self.policy = policy_model
def evaluate_judge_accuracy(
self,
test_problems: List[Dict],
max_problems: int = 500
) -> Dict:
"""
Evaluate judge accuracy on intermediate steps.
Metrics:
- Accuracy: % of steps correctly classified
- F1: balance between precision and recall
"""
correct_verdicts = 0
total_verdicts = 0
tp, fp, fn = 0, 0, 0
for problem_data in test_problems[:max_problems]:
problem = problem_data["problem"]
ground_truth_steps = problem_data["steps"]
correct_step_labels = problem_data["correct_steps"]
for step_idx, step in enumerate(ground_truth_steps):
output = self.judge(
problem,
ground_truth_steps[:step_idx],
step
)
pred_is_correct = output["verdict_logits"].argmax().item() == 1
is_correct = correct_step_labels[step_idx]
if pred_is_correct == is_correct:
correct_verdicts += 1
if is_correct:
tp +=
:
pred_is_correct:
fp +=
:
fn +=
total_verdicts +=
accuracy = correct_verdicts / total_verdicts
precision = tp / (tp + fp) (tp + fp) >
recall = tp / (tp + fn) (tp + fn) >
f1 = * (precision * recall) / (precision + recall) (precision + recall) >
{
: accuracy,
: precision,
: recall,
: f1
}
() -> :
results = {
: ,
: ,
:
}
problem_data test_problems:
problem = problem_data[]
target = problem_data[]
solution_no_judge = .policy.generate_solution(problem)
correct_no_judge = .check_solution(solution_no_judge, target)
reasoner = ReasoningWithStepValidation(.policy, .judge)
output = reasoner.generate_with_validation(problem)
solution_with_judge = .join(output[])
correct_with_judge = .check_solution(solution_with_judge, target)
correct_no_judge:
results[] +=
correct_with_judge:
results[] +=
total = (test_problems)
results[] /= total
results[] /= total
results[] = results[] - results[]
results
() -> :
solution.strip() == target.strip()
Practical Guidance
Training Recipe
- Phase 1: Train judge on synthetic step labels (easy supervision)
- Phase 2: RL training on relative outcome comparisons (harder data)
- Phase 3: Fine-tune on domain-specific problems
Explanation Integration
- Explanations are emergent from RL training (not explicitly supervised)
- Can optionally supervise with human-written explanations
- Improves interpretability without hurting performance
When to Use
- Reasoning tasks requiring intermediate step validation
- Multi-step math and coding problems
- Scenarios where explainability is important
- Improving reasoning through search (beam search, etc.)
When NOT to Use
- Single-step tasks without intermediate verification
- Real-time systems (judge adds latency)
- Domains without clear step correctness
Performance Expectations
- Judge accuracy: 85-92% on intermediate steps
- Reasoning improvement: +3-8% on complex problems
- Inference overhead: 1.5-2x slowdown (due to validation)
Reference
StepWiser: Stepwise Generative Judges for Wiser Reasoning. arXiv:2508.19229