| name | reinforcement-pretraining |
| title | Reinforcement Pre-Training: RL-Based Scaling for LLM Foundation Development |
| version | 0.0.2 |
| engine | skillxiv-v0.0.2-claude-opus-4.6 |
| license | MIT |
| url | https://arxiv.org/abs/2506.08007 |
| keywords | ["reinforcement-learning","pretraining","scaling-laws","next-token-prediction","reasoning"] |
| description | Apply reinforcement learning to pre-training by framing next-token prediction as a reasoning task with verifiable rewards, achieving superior scaling compared to standard language modeling. |
Reinforcement Pre-Training: RL-Based Scaling for LLM Foundation Development
Core Concept
Reinforcement Pre-Training (RPT) reframes next-token prediction as a reasoning task trained with reinforcement learning rather than supervised learning. Instead of matching target tokens directly, models generate reasoning chains and receive binary rewards for correctly predicting the next token. This approach leverages vast text corpora as a naturally verifiable reward signal, enabling general-purpose RL during pre-training without external annotations. RPT achieves stronger scaling properties and matches larger supervised baselines using significantly fewer training tokens.
Architecture Overview
- Reasoning-Capable Foundation: Built on models like Deepseek-R1-Distill-Qwen-14B with chain-of-thought capability
- Verifiable Reward System: Binary rewards based on exact token prediction, eliminating reward hacking
- On-Policy GRPO Training: Group Relative Policy Optimization with typically 8 parallel rollouts
- Entropy-Based Data Filtering: Prioritizes challenging token positions where reasoning matters most
- Scaling Law Alignment: Consistent power-law improvement with compute, following predictable patterns
Implementation
Step 1: Set Up Reward Function and Filtering
Implement the core verifiable reward mechanism:
import torch
import numpy as np
from transformers import AutoTokenizer
class NextTokenRewardFunction:
"""Verifiable reward based on next-token prediction accuracy"""
def __init__(self, tokenizer):
self.tokenizer = tokenizer
def compute_reward(self, generated_text, ground_truth_text, prediction_mode='prefix'):
"""
Compute binary reward for next-token prediction.
Args:
generated_text: model's generated reasoning + prediction
ground_truth_text: ground truth completion
prediction_mode: 'prefix' for exact prefix matching
Returns:
reward: 1.0 if generated is exact prefix, 0.0 otherwise
"""
if prediction_mode == 'prefix':
is_prefix = ground_truth_text.startswith(generated_text)
return float(is_prefix)
return 0.0
def entropy_based_filtering(texts, tokenizer, entropy_threshold=2.0):
"""
Filter to include only challenging token positions.
Uses Shannon entropy to identify positions where prediction is non-trivial.
"""
filtered_texts = []
for text in texts:
tokens = tokenizer.encode(text)
entropies = []
for t in range(1, len(tokens)):
prev_tokens = tokens[:t]
next_token = tokens[t]
entropy = estimate_entropy(prev_tokens, next_token)
entropies.append(entropy)
challenging_positions = [i i, e (entropies) e > entropy_threshold]
(challenging_positions) > :
filtered_texts.append({
: text,
: challenging_positions,
: np.mean(entropies)
})
filtered_texts
():
math
math.log(vocab_size) *
Step 2: Implement RPT Training Loop with GRPO
Create the reinforcement pre-training pipeline:
import torch.nn.functional as F
from torch.optim import AdamW
def rpt_training_step(model, prefix_ids, ground_truth_ids, reward_fn, group_size=8,
learning_rate=1e-6):
"""
Single RPT training step using GRPO.
Args:
model: Language model
prefix_ids: Input context
ground_truth_ids: Ground truth next tokens
reward_fn: Reward function instance
group_size: Number of parallel rollouts (G=8 typical)
"""
batch_size = prefix_ids.shape[0]
device = prefix_ids.device
rollout_ids = []
rollout_log_probs = []
rewards = torch.zeros(batch_size, group_size, device=device)
model.eval()
with torch.no_grad():
for g in range(group_size):
generated = generate_with_reasoning(model, prefix_ids, max_new_tokens=128)
rollout_ids.append(generated)
for b in range(batch_size):
generated_text = model.tokenizer.decode(generated[b])
ground_truth_text = model.tokenizer.decode(ground_truth_ids[b])
reward = reward_fn.compute_reward(generated_text, ground_truth_text)
rewards[b, g] = reward
model.train()
optimizer = AdamW(model.parameters(), lr=learning_rate)
for g in range(group_size):
optimizer.zero_grad()
generated = rollout_ids[g]
outputs = model(input_ids=generated)
logits = outputs.logits
log_probs = torch.log_softmax(logits[:, :-, :], dim=-)
action_log_probs = log_probs.gather(, generated[:, :].unsqueeze()).squeeze()
sequence_log_prob = action_log_probs.(dim=)
group_reward = rewards[:, g]
baseline = rewards.mean(dim=)
advantage = group_reward - baseline
loss = -(sequence_log_prob * advantage.detach()).mean()
loss.backward()
torch.nn.utils.clip_grad_norm_(model.parameters(), )
optimizer.step()
{
: loss.item(),
: rewards.mean().item(),
: rewards.std().item()
}
():
device = prefix_ids.device
reasoning_ids = model.generate(
input_ids=prefix_ids,
max_new_tokens=(max_new_tokens, ),
temperature=temperature,
do_sample=,
pad_token_id=model.config.eos_token_id
)
output_ids = model.generate(
input_ids=reasoning_ids,
max_new_tokens=(, max_new_tokens - ),
temperature=temperature,
do_sample=,
pad_token_id=model.config.eos_token_id
)
output_ids
Step 3: Implement Scaling Study Infrastructure
Track scaling laws across different training scales:
import matplotlib.pyplot as plt
from scipy.optimize import curve_fit
class ScalingStudy:
def __init__(self, model_size=14):
self.model_size = model_size
self.results = []
def power_law(self, x, a, b):
"""Power law: f(x) = a * x^(-b)"""
return a * (x ** (-b))
def run_scaling_experiment(self, token_budgets=[1e8, 5e8, 1e9, 5e9, 1e10]):
"""Run RPT with different compute budgets"""
for tokens in token_budgets:
print(f"Training with {tokens:.1e} tokens...")
model = load_base_model()
train_loss = train_rpt_model(model, num_tokens=int(tokens))
superglue_score = evaluate_superglue(model)
mmlu_pro_score = evaluate_mmlu_pro(model)
self.results.append({
'tokens': tokens,
'loss': train_loss,
'superglue': superglue_score,
'mmlu_pro': mmlu_pro_score
})
return self.results
def ():
tokens = [r[] r .results]
mmlu_scores = [r[] r .results]
popt, _ = curve_fit(.power_law, tokens, mmlu_scores, p0=[, ])
plt.figure(figsize=(, ))
plt.loglog(tokens, mmlu_scores, , label=)
plt.loglog(tokens, .power_law(np.array(tokens), *popt), ,
label=)
plt.xlabel()
plt.ylabel()
plt.legend()
plt.savefig()
popt
Step 4: Comparative Evaluation
Compare RPT against supervised baselines:
def comparative_evaluation(rpt_model, supervised_baseline, test_datasets):
"""
Compare RPT vs. supervised baselines.
RPT typically matches 2x+ larger supervised models using fewer tokens.
"""
results = {
'rpt': {},
'supervised': {}
}
for dataset_name, dataset in test_datasets.items():
rpt_score = evaluate_dataset(rpt_model, dataset)
results['rpt'][dataset_name] = rpt_score
supervised_score = evaluate_dataset(supervised_baseline, dataset)
results['supervised'][dataset_name] = supervised_score
improvement = ((rpt_score - supervised_score) / supervised_score) * 100
print(f"{dataset_name}: RPT={rpt_score:.2%}, Supervised={supervised_score:.2%}, "
f"Improvement={improvement:+.1f}%")
return results
Practical Guidance
- Base Model Choice: Use reasoning-capable models (Deepseek-R1, DeepSeek-V3) as foundation
- Sequence Length: Use 8k token sequences; longer contexts reduce FLOPs per token but increase memory
- Learning Rate: Start with 1×10⁻⁶; scale down for larger models to prevent instability
- Group Size: G=8 provides good balance between gradient quality and computational cost
- Entropy Filtering: Prioritize hard examples (high entropy positions) for stable training
- Reward Design: Binary prefix matching is stable; avoid continuous rewards that encourage reward hacking
- Data Selection: Works best on well-curated, diverse text with challenging predictions
- Downstream RL: RPT models serve as excellent foundation for further RL fine-tuning
Reference
- RPT achieves better scaling by converting supervised loss to RL problem with natural verifiability
- Next-token prediction provides abundant, cost-free reward signal without external annotation
- Reasoning chains enable models to solve harder predictions before committing to token choice
- Scaling laws consistently follow power-law patterns with good fit, enabling predictability