| name | revisual-multimodal-reasoning |
| title | ReVisual-R1: Multimodal Reasoning with Cold-Start and Staged RL |
| version | 0.0.2 |
| engine | skillxiv-v0.0.2-claude-opus-4.6 |
| license | MIT |
| url | https://arxiv.org/abs/2506.04207 |
| keywords | ["multimodal-reasoning","reinforcement-learning","curriculum-learning","vision-language"] |
| description | Develop sophisticated multimodal reasoning through text-centric cold-start initialization, prioritized advantage distillation, and staged RL refinement. |
ReVisual-R1: Advancing Multimodal Reasoning
Core Concept
ReVisual-R1 demonstrates that unlocking multimodal reasoning capabilities requires a carefully designed three-stage curriculum: text-only cold-start with complex reasoning examples, multimodal RL with prioritized advantage distillation to prevent gradient stagnation, and text-only refinement to consolidate linguistic fluency. This approach outperforms simultaneous mixed-modality training and achieves state-of-the-art performance among 3B/7B open-source models.
Architecture Overview
- Stage 1: Text-Centric Cold Start - 283K curated high-difficulty reasoning examples (language-only)
- Stage 2: Multimodal RL with PAD - GRPO enhanced by Prioritized Advantage Distillation
- Stage 3: Text-Only RL Refinement - Polish linguistic quality and reasoning consistency
- Prioritized Advantage Distillation (PAD): Filters zero-advantage samples and resamples informative trajectories
- Challenge Addressed: Gradient stagnation in multimodal GRPO settings where standard reward signals provide weak gradients
Implementation
Step 1: Prepare Cold-Start Text-Only Dataset
import json
from typing import List, Dict
class ColdStartDatasetBuilder:
def __init__(self, num_samples=283000):
self.target_size = num_samples
self.difficulty_threshold = "high"
def collect_complex_reasoning_examples(self):
"""Gather high-difficulty text reasoning from multiple sources"""
sources = {
'math_olympiad': self.source_olympiad_problems(),
'theoretical_physics': self.source_physics_proofs(),
'algorithm_design': self.source_algorithm_challenges(),
'logical_deduction': self.source_logic_puzzles(),
}
dataset = []
for source_name, examples in sources.items():
for example in examples:
if example['complexity'] == 'high':
dataset.append({
'question': example['prompt'],
'reasoning_chain': example['solution'],
'source': source_name,
'difficulty_score': example['difficulty'],
'reasoning_depth': example['step_count']
})
dataset[:.target_size]
():
augmented = []
example examples:
structured_chain = .reformat_chain_of_thought(
example[]
)
augmented.append({
**example,
: structured_chain
})
augmented
():
steps = reasoning_text.split()
formatted = []
i, step (steps, ):
formatted.append()
.join(formatted)
():
builder = ColdStartDatasetBuilder(num_samples=)
cold_start_data = builder.collect_complex_reasoning_examples()
cold_start_data = builder.augment_with_synthetic_reasoning(cold_start_data)
Step 2: Implement Prioritized Advantage Distillation
import torch
import numpy as np
class PrioritizedAdvantageDistillation:
def __init__(self, temperature=1.0):
self.temperature = temperature
def compute_advantages(self, rewards, baseline_values):
"""Calculate advantage for each trajectory"""
advantages = rewards - baseline_values
mean_adv = np.mean(advantages)
std_adv = np.std(advantages)
normalized_advantages = (advantages - mean_adv) / (std_adv + 1e-8)
return normalized_advantages
def filter_zero_advantage_samples(self, trajectories, advantages,
threshold=0.01):
"""Remove non-informative samples where advantage ≈ 0"""
filtered = []
valid_indices = []
for idx, (traj, adv) in enumerate(zip(trajectories, advantages)):
if abs(adv) > threshold:
filtered.append(traj)
valid_indices.append(idx)
print(f"Filtered {len(trajectories) - len(filtered)} zero-advantage samples")
print(f"Retained {len(filtered)}/{len(trajectories)} informative samples")
return filtered, valid_indices
():
exp_advantages = np.exp(advantages / temperature)
probabilities = exp_advantages / np.(exp_advantages)
num_samples = (trajectories)
resampled_indices = np.random.choice(
num_samples,
size=num_samples,
p=probabilities,
replace=
)
resampled = [trajectories[i] i resampled_indices]
resampled, resampled_indices
:
():
.model = model
.pad = pad_module
():
responses_text = .model.generate(batch_text)
responses_mm = .model.generate(batch_multimodal)
rewards_text = [reward_fn(r) r responses_text]
rewards_mm = [reward_fn(r) r responses_mm]
baseline_text = .estimate_baseline(batch_text)
baseline_mm = .estimate_baseline(batch_multimodal)
advantages_text = .pad.compute_advantages(
np.array(rewards_text),
baseline_text
)
advantages_mm = .pad.compute_advantages(
np.array(rewards_mm),
baseline_mm
)
responses_filtered, _ = .pad.filter_zero_advantage_samples(
responses_mm, advantages_mm
)
responses_resampled, _ = .pad.prioritized_resampling(
advantages_mm[_], responses_mm, temperature=
)
policy_loss = .compute_policy_loss(
responses_resampled,
advantages_mm
)
optimizer.zero_grad()
policy_loss.backward()
optimizer.step()
policy_loss.item(), (responses_filtered), (responses_resampled)
Step 3: Implement Three-Stage Training Pipeline
class ThreeStageTrainingPipeline:
def __init__(self, model, cold_start_data):
self.model = model
self.cold_start_data = cold_start_data
self.pad = PrioritizedAdvantageDistillation()
def stage_1_text_coldstart(self, num_epochs=3):
"""
Train exclusively on text-only complex reasoning.
Duration: Initialize model reasoning capability.
"""
print("=== Stage 1: Text-Only Cold Start ===")
print(f"Dataset size: {len(self.cold_start_data)} examples")
print(f"Focus: Complex reasoning chains, logical coherence")
optimizer = self.get_optimizer(lr=1e-4)
for epoch in range(num_epochs):
total_loss = 0
for batch in self.create_batches(self.cold_start_data, bs=32):
loss = self.model.compute_language_loss(
batch['question'],
batch['structured_reasoning']
)
optimizer.zero_grad()
loss.backward()
optimizer.step()
total_loss += loss.item()
print(f"Epoch {epoch+1}: Loss = ")
():
()
()
()
optimizer = .get_optimizer(lr=)
iteration (num_iterations):
batch = .sample_multimodal_batch(multimodal_data, bs=)
loss, filtered_count, resampled_count = (
GRPOWithPAD(.model, .pad).multimodal_grpo_step(
batch[],
batch[],
reward_fn=.compute_multimodal_reward,
optimizer=optimizer
)
)
(iteration + ) % == :
(
)
():
()
()
()
optimizer = .get_optimizer(lr=)
epoch (num_epochs):
batch .create_batches(text_data, bs=):
text_output = .model.generate(batch[])
fluency_reward = .evaluate_fluency(text_output)
coherence_reward = .evaluate_coherence(text_output)
combined_reward = * fluency_reward + * coherence_reward
log_probs = .model.log_probability(text_output)
loss = -log_probs * (combined_reward - )
optimizer.zero_grad()
loss.backward()
optimizer.step()
Practical Guidance
-
Text-First Initialization: Begin with 283K+ high-difficulty text-only examples before introducing visual data. This establishes reasoning foundations that transfer to multimodal tasks.
-
Detect Gradient Stagnation: Monitor loss plateaus during multimodal GRPO training. If loss stops decreasing despite correct implementation, PAD is likely needed.
-
PAD Implementation Details: Filter out advantages near zero (threshold ≈ 0.01), then resample using temperature-controlled softmax distribution over remaining samples. This concentrates training on informative trajectories.
-
Curriculum Progression: Don't skip stages or overlap them. The three-stage design ensures: (1) reasoning foundation, (2) multimodal alignment, (3) linguistic polish. Simultaneous mixed training fails empirically.
-
Evaluation Protocol: Test performance across text-only, vision-only, and vision-language tasks. Quality improvements should generalize across modalities, indicating robust reasoning rather than overfitting to training distribution.
Reference
- Paper: Advancing Multimodal Reasoning (2506.04207)
- Cold Start: 283K curated complex reasoning examples
- Method: GRPO + Prioritized Advantage Distillation
- Results: State-of-the-art for 3B/7B open-source models on multimodal reasoning