| name | vla-r1-reasoning-embodied-agents |
| title | VLA-R1: Enhancing Reasoning in Vision-Language-Action Models |
| version | 0.0.2 |
| engine | skillxiv-v0.0.2-claude-opus-4.6 |
| license | MIT |
| url | https://arxiv.org/abs/2510.01623 |
| keywords | ["VLA","reasoning","embodied-AI","RLVR","robotics","chain-of-thought"] |
| description | Enhance embodied robot reasoning by integrating explicit chain-of-thought supervision with reinforcement learning from verifiable rewards (GRPO+RL). Use when improving robot decision-making for tasks requiring spatial reasoning and constraint satisfaction. |
VLA-R1: Enhancing Reasoning in Vision-Language-Action Models
VLA-R1 addresses a critical gap in Vision-Language-Action models: they emit final actions directly without explicit reasoning over affordances, geometric relations, or constraints. By combining chain-of-thought supervision with RL from verifiable rewards, the model learns to reason about the task before acting.
Core Architecture
- Chain-of-Thought annotation: 13,000 examples with explicit reasoning steps (VLA-CoT-13K dataset)
- Supervised fine-tuning phase: Train reasoning head alongside action prediction
- Reinforcement learning phase: Optimize reasoning quality via GRPO with verifiable geometric rewards
- Multimodal grounding: Maintains visual understanding while developing reasoning capability
Implementation Steps
Setup reasoning-enhanced VLA framework:
from vla_r1 import ReasoningVLATrainer, VLACoTDataset
vla_model = ReasoningVLATrainer(
base_model="Qwen2.5-VL-7B",
include_reasoning_head=True,
reasoning_budget=256,
action_head_architecture="standard"
)
dataset = VLACoTDataset(
num_examples=13000,
domains=["tabletop_manipulation", "grasping", "placement"],
reasoning_style="explicit_affordance"
)
Execute SFT phase with reasoning supervision:
sft_trainer = vla_model.create_sft_trainer(
learning_rate=1e-4,
batch_size=8,
num_epochs=2
)
sft_losses = []
for batch in dataset:
reasoning_output = vla_model.generate_reasoning(batch["image"])
action_logits = vla_model.predict_action(
image=batch[],
reasoning=reasoning_output,
instruction=batch[]
)
reasoning_loss = compute_language_loss(
predictions=reasoning_output,
targets=batch[]
)
action_loss = compute_action_loss(
predictions=action_logits,
targets=batch[]
)
total_loss = reasoning_loss + action_loss
total_loss.backward()
sft_losses.append(total_loss.item())