| name | bro-rl-broad-rollout-scaling |
| title | BroRL: Scaling via Broad Exploration Rather Than Longer Training |
| version | 0.0.2 |
| engine | skillxiv-v0.0.2-claude-opus-4.6 |
| license | MIT |
| url | https://arxiv.org/abs/2510.01180 |
| keywords | ["RLVR","scaling","exploration","reasoning","efficiency"] |
| description | Overcome reasoning model training plateaus by increasing rollouts per prompt (N=512) rather than training steps, addressing unsampled coupling that destabilizes learning. Theoretical analysis shows broad exploration eliminates plateau bottleneck. |
BroRL: Scaling via Broad Exploration Rather Than Longer Training
BroRL addresses a fundamental plateau problem in reasoning model RL: training stops improving after ~3,000 steps. Rather than longer training, the solution is broader exploration. By increasing rollouts per prompt from 16 to 512, models escape the plateau and continue improving, grounded in theoretical analysis of unsampled coupling destabilization.
Core Architecture
- Broad exploration principle: Increase N (rollouts per prompt) instead of training steps
- Theoretical foundation: Unsampled coupling term analysis explains plateau mechanism
- Scaling formula: Learning rate adjustments based on N following principled schedule
- Memory-bound to compute-bound: Shift from GPU memory bottleneck to compute bottleneck
Implementation Steps
Configure BroRL rollout and learning rate strategy:
from bro_rl import BroadRLTrainer, RolloutScaler
trainer = BroadRLTrainer(
model=your_reasoning_llm,
base_rollouts=16,
target_rollouts=512,
algorithm="GRPO"
)
rollout_scaler = RolloutScaler(
base_learning_rate=1e-5,
rollout_schedule=[16, 32, 64, 128, 256, 512]
)
learning_rates = rollout_scaler.compute_schedule()
Execute broad rollout training:
training_steps = 3000
epochs_per_rollout_config = 500
for config_idx, (rollout_count, learning_rate) (
(rollout_scaler.rollout_schedule, learning_rates)
):
()
optimizer = torch.optim.AdamW(
model.parameters(),
lr=learning_rate,
betas=(, ),
weight_decay=
)
step (epochs_per_rollout_config):
batch training_dataloader:
prompts = batch[]
rollouts = []
_ (rollout_count):
rollout = model.generate(
prompts,
max_length=,
temperature=,
top_p=
)
rollouts.append(rollout)
rewards = verifier.evaluate_batch(rollouts)
advantages = compute_advantages(
rewards=rewards,
baseline_method=
)
log_probs = model.compute_log_probs(rollouts)
policy_loss = -((log_probs * advantages).(dim=)).mean()
kl_loss = compute_kl_divergence(
model=model,
reference_model=reference_model,
kl_weight=
)
total_loss = policy_loss + kl_loss
optimizer.zero_grad()
total_loss.backward()
torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=)
optimizer.step()
step % == :
reference_model.load_state_dict(model.state_dict())
step % == :
success_rate = (rewards > ).().mean()
(
)