| name | proactive-self-refinement |
| title | A Stitch in Time: Proactive Self-Refinement for Language Models |
| version | 0.0.2 |
| engine | skillxiv-v0.0.2-claude-opus-4.6 |
| license | MIT |
| url | https://arxiv.org/abs/2508.12903 |
| keywords | ["self-refinement","dynamic-revision","in-generation-improvement","token-efficiency"] |
| description | Enable models to refine outputs dynamically during generation based on internal signals, reducing token consumption by 41.6% while improving accuracy by 8.2%. |
A Stitch in Time: Proactive Self-Refinement for Language Models
Core Concept
Traditional self-refinement works in fixed cycles: generate → evaluate → regenerate. This is inefficient because the model can't start refinement until generation completes.
Proactive Active Self-Refinement (PASR) lets models decide dynamically during generation whether, when, and how to refine. The model learns to detect when its reasoning is going wrong and self-correct mid-generation, like humans revising thoughts while speaking.
Architecture Overview
- Dynamic Refinement Trigger: Learn to detect when refinement is needed
- In-Generation Revision: Backtrack and revise during generation
- Internal Quality Signals: Use model's own uncertainty/confidence
- Learned Refinement Strategy: Decide how aggressively to refine
- Token Efficiency: Avoid wasteful regeneration of correct portions
- Adaptive Refinement: Different tasks get different refinement patterns
Implementation Steps
1. Define Refinement Points and Signals
import torch
import torch.nn as nn
from typing import List, Tuple
class RefinementSignal:
"""Detect when refinement is needed"""
def __init__(self, model):
self.model = model
def compute_confidence(self, logits: torch.Tensor) -> float:
"""Compute model's confidence in current prediction"""
probs = torch.softmax(logits, dim=-1)
confidence = probs.max().item()
return confidence
def compute_entropy(self, logits: torch.Tensor) -> float:
"""Compute entropy of output distribution"""
probs = torch.softmax(logits, dim=-1)
entropy = -(probs * torch.log(probs + 1e-10)).sum().item()
return entropy
def compute_consistency(self, logits_list: List[torch.Tensor]) -> float:
"""Compute consistency across multiple forward passes"""
if len(logits_list) < 2:
return 1.0
predictions = [logits.argmax().item() logits logits_list]
consistency = predictions.count(predictions[]) / (predictions)
consistency
() -> :
low_confidence = confidence <
high_entropy = entropy >
reasonable_length = token_count < max_tokens *
low_confidence high_entropy reasonable_length
2. Implement Dynamic Backtracking
class DynamicBacktracker:
"""Backtrack and revise during generation"""
def __init__(self, model, tokenizer):
self.model = model
self.tokenizer = tokenizer
self.token_history = []
def find_refinement_point(self, current_tokens: List[int],
quality_scores: List[float]) -> int:
"""Find best point to backtrack to"""
min_quality_idx = 0
min_quality = quality_scores[0]
for i, score in enumerate(quality_scores):
if score < min_quality:
min_quality = score
min_quality_idx = i
min_backtrack = int(len(current_tokens) * 0.2)
backtrack_point = max(min_backtrack_idx, min_backtrack)
return backtrack_point
def revise_from_point(self, context_tokens: List[int],
backtrack_point: int, num_alternatives: int = 3):
"""Generate alternatives from backtrack point"""
revised_tokens = context_tokens[:backtrack_point]
alternatives = []
temp [, , ]:
output = .model.generate(
torch.tensor([revised_tokens]),
max_new_tokens=,
temperature=temp,
do_sample=
)
alternatives.append(output[].tolist())
alternatives
() -> [[], ]:
best_seq =
best_quality = -
alt alternatives:
text = .tokenizer.decode(alt)
quality = quality_fn(text)
quality > best_quality:
best_quality = quality
best_seq = alt
best_seq, best_quality
3. Train Refinement Policy
class RefinementPolicy(nn.Module):
"""Learn when and how to refine"""
def __init__(self, hidden_size=768):
super().__init__()
self.encoder = nn.TransformerEncoderLayer(
d_model=hidden_size,
nhead=8,
dim_feedforward=2048,
batch_first=True
)
self.refinement_trigger = nn.Linear(hidden_size, 1)
self.backtrack_distance = nn.Linear(hidden_size, 100)
self.refinement_intensity = nn.Linear(hidden_size, 1)
def forward(self, token_embeddings: torch.Tensor) -> dict:
"""
Decide refinement strategy
Args:
token_embeddings: [seq_len, hidden_size] embeddings of generated tokens
Returns:
refinement_decision: dict with trigger, backtrack_distance, intensity
"""
context = self.encoder(token_embeddings.unsqueeze(0))
context = context[0, -1, :]
refine_logit = self.refinement_trigger(context)
refine_prob = torch.sigmoid(refine_logit)
backtrack_logits = .backtrack_distance(context)
backtrack_dist = torch.softmax(backtrack_logits, dim=)
intensity = torch.sigmoid(.refinement_intensity(context))
{
: refine_prob.item() > ,
: refine_prob.item(),
: backtrack_dist.argmax().item(),
: intensity.item()
}
():
optimizer = torch.optim.Adam(policy.parameters(), lr=)
epoch (num_epochs):
batch train_data:
prompts = batch[]
target_outputs = batch[]
losses = []
prompt, target (prompts, target_outputs):
generated_tokens = []
token_embeddings = []
refinements_applied =
tokens = tokenizer.encode(prompt)
step ():
embeddings = model.get_embeddings(torch.tensor([tokens]))
decision = policy(embeddings[])
token_embeddings.append(embeddings[, -, :])
decision[] step > :
backtrack_dist = decision[]
tokens = tokens[:-backtrack_dist]
refinements_applied +=
logits = model(torch.tensor([tokens])).logits[, -, :]
next_token = logits.argmax().item()
tokens.append(next_token)
next_token == tokenizer.eos_token_id:
generated_text = tokenizer.decode(tokens)
similarity = compute_similarity(generated_text, target)
efficiency_bonus = ( - (tokens) / )
reward = * similarity + * efficiency_bonus
loss = -torch.tensor(reward)
losses.append(loss)
total_loss = torch.stack(losses).mean()
optimizer.zero_grad()
total_loss.backward()
optimizer.step()
4. Inference with PASR
def generate_with_pasr(model, policy, tokenizer, prompt: str,
max_length: int = 256, refinement_threshold: float = 0.5):
"""Generate with dynamic proactive self-refinement"""
tokens = tokenizer.encode(prompt)
generated = []
refinement_count = 0
quality_scores = []
for step in range(max_length):
embeddings = model.get_embeddings(torch.tensor([tokens]))
curr_embedding = embeddings[0, -1, :]
decision = policy(curr_embedding.unsqueeze(0))
confidence = decision['refine_probability']
quality_scores.append(confidence)
if (decision['should_refine'] and
confidence < refinement_threshold and
step > 10 and len(generated) > 5):
backtrack = decision['backtrack_distance']
tokens = tokens[:-min(backtrack, len(generated))]
generated = generated[:-min(backtrack, len(generated))]
refinement_count += 1
logits = model(torch.tensor([tokens])).logits[0, -1, :]
logits = logits /
next_token = torch.multinomial(
torch.softmax(logits, dim=-), num_samples=
).item()
:
logits = model(torch.tensor([tokens])).logits[, -, :]
next_token = logits.argmax().item()
tokens.append(next_token)
generated.append(next_token)
next_token == tokenizer.eos_token_id:
result_text = tokenizer.decode(generated)
{
: result_text,
: (generated),
: refinement_count,
: - (refinement_count / (generated))
}
5. Evaluation
def evaluate_pasr(model, policy, tokenizer, benchmark_tasks):
"""Evaluate PASR on accuracy and efficiency"""
accuracy = 0.0
token_efficiency = 0.0
num_tasks = 0
for task in benchmark_tasks:
prompt = task['prompt']
target = task['target']
result = generate_with_pasr(model, policy, tokenizer, prompt)
generated_text = result['text']
is_correct = check_correctness(generated_text, target)
accuracy += 1.0 if is_correct else 0.0
baseline_tokens = 100
token_efficiency += 1.0 - (result['tokens_generated'] / baseline_tokens)
num_tasks += 1
avg_accuracy = accuracy / num_tasks if num_tasks > 0 else 0.0
avg_efficiency = token_efficiency / num_tasks if num_tasks > 0 else 0.0
print(f"Accuracy: {avg_accuracy * 100:.1f}%")
print(f"Token Efficiency: {avg_efficiency * 100:.1f}%")
return avg_accuracy, avg_efficiency
Practical Guidance
- Refinement Threshold: 0.4-0.6 (lower = more aggressive refinement)
- Backtrack Distance: 5-20 tokens (avoid over-revision)
- Temperature: 0.7-0.9 for alternatives (higher = more diversity)
- Policy Training: Mix supervised + RL (80% supervised, 20% RL)
- Quality Function: Use task-specific metrics (BLEU, exact match, etc.)
Reference
A Stitch in Time (2508.12903): https://arxiv.org/abs/2508.12903
Enable dynamic, in-generation self-refinement where models decide when to backtrack and revise, achieving 41.6% token reduction and 8.2% accuracy improvement over baseline generation.