| name | critique-rl-training |
| title | Critique-RL: Training Language Models for Critiquing through Two-Stage RL |
| version | 0.0.2 |
| engine | skillxiv-v0.0.2-claude-opus-4.6 |
| license | MIT |
| url | https://arxiv.org/abs/2510.24320 |
| keywords | ["RL","Feedback Learning","Critic Models","Two-stage Training","Reasoning"] |
| description | Trains language models to provide quality feedback through two-stage RL. Stage 1 optimizes discriminability (distinguishing good vs bad responses). Stage 2 adds helpfulness rewards (improving actor after feedback). Achieves 9.02% improvement without requiring stronger supervisors for training data. |
Critique-RL: Learning to Provide Effective Feedback
Standard RL for critic models focuses on generating feedback, but doesn't ensure feedback quality. Critique-RL uses two-stage training to develop both discriminative ability (telling good from bad) and helpfulness (guiding improvement).
The approach enables cheaper critic training without external supervisors.
Core Concept
Two-stage reinforcement learning:
- Stage 1: Optimize discriminability via direct reward signals
- Stage 2: Optimize helpfulness via actor improvement signals
- Stage 1 prevents feedback collapse into neutral comments
- Stage 2 ensures feedback actually helps (no false positives)
Architecture Overview
- Actor-Critic pair: actor generates responses, critic evaluates
- Rule-based reward for Stage 1 (quality assessment)
- Actor improvement reward for Stage 2 (helpfulness)
- Regularization to preserve Stage 1 discriminability
Implementation Steps
Implement the two-stage training pipeline with reward signals:
class TwoStageCriticRL:
def __init__(self, critic_model, actor_model):
self.critic = critic_model
self.actor = actor_model
self.critic_optimizer = torch.optim.AdamW(critic_model.parameters())
self.stage = 'discriminability'
def stage1_discriminability(self, good_responses, bad_responses):
"""Stage 1: Train to distinguish quality levels."""
for good, bad in zip(good_responses, bad_responses):
good_score = self.critic(good)['quality']
bad_score = self.critic(bad)['quality']
loss = -torch.log(torch.sigmoid(good_score - bad_score))
.critic_optimizer.zero_grad()
loss.backward()
.critic_optimizer.step()
():
prompt, initial_response actor_feedback_pairs:
feedback = .critic(initial_response)[]
improved = ._improve_via_feedback(
initial_response, feedback, num_steps
)
initial_score = ._evaluate_response(initial_response)
improved_score = ._evaluate_response(improved)
improvement = improved_score - initial_score
improvement > :
feedback_score = .critic(initial_response)[]
rl_loss = -feedback_score * improvement
.critic_optimizer.zero_grad()
rl_loss.backward()
.critic_optimizer.step()
._discriminability_regularization()
():
current = response
_ (num_steps):
improved = .actor(current + + feedback)
current = improved
current
():
.critic(response)[]
():
validation_samples = []
good, bad validation_samples:
good_score = .critic(good)[]
bad_score = .critic(bad)[]
good_score < bad_score:
penalty = torch.nn.functional.relu(bad_score - good_score)
.critic_optimizer.zero_grad()
penalty.backward()
.critic_optimizer.step()