Skip to content

Folders and files

NameName
Last commit message
Last commit date

Latest commit

 

History

28 Commits
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 

Repository files navigation

Ego to World: Collaborative Spatial Reasoning in Embodied Systems via Reinforcement Learning

arXiv Project Page Python 3.8+ License Models Dataset

A framework for cross-view spatial reasoning with vision-language models. CoRL trains VLMs using Supervised Fine-Tuning (SFT) followed by Reinforcement Learning with Cross-View Spatial Reward (CVSR) for improved spatial understanding in dual-viewpoint robotic scenarios.

Why cross-view reasoning? Single-view reasoning fails when objects are occluded — cross-view reasoning uses overlapping views as references to correctly locate and grasp targets.

Examples

The CoRL model performs cross-view spatial reasoning with chain-of-thought. Below are outputs from Qwen2.5-VL-3B-CoRL:

E2W-1: Object Counting

View A (main) View B (auxiliary)

Q: Can you assess how many apple are present?

Model: <think> I can see a apple at [195, 74, 207, 85] in View A ... a apple at [111, 129, 124, 139] in View B ... The apple in auxiliary perspective is not visible in main perspective ... Combining observations from both perspectives, the total is 2. </think> \boxed{2}

E2W-2: Spatial Reasoning

View A (main) View B (auxiliary)

Q: Which object is the third from the left to right on the table?

Model: <think> In View B: cardboardbox at [4,93,49,107], pizza at [136,66,159,87], kettle at [204,83,226,103], hammer at [328,94,360,106]. Spatial order left→right: cardboardbox → pizza → kettle → hammer. The third is kettle. </think> \boxed{kettle}

E2W-3: Grasping

View A (main) View B (auxiliary)

Q: Please grasp the apple the third from the front to back on the table.

Model: <think> The target apple is visible in main perspective at [248, 89, 255, 95]. Center coordinates: [251.5, 92.0]. I should grasp at image 0 coordinates [251.5, 92.0]. </think> \boxed{0, [251.5, 92.0]}

Quick Start

Inference with Pre-trained Models

pip install transformers torch Pillow

python scripts/inference.py \
    --model henggg/Qwen2.5-VL-3B-CoRL \
    --images view_a.jpg view_b.jpg \
    --question "How many objects are on the table?"

Or in Python:

from transformers import AutoModelForVision2Seq, AutoProcessor
from PIL import Image

model = AutoModelForVision2Seq.from_pretrained(
    "henggg/Qwen2.5-VL-3B-CoRL",
    torch_dtype="bfloat16",
    attn_implementation="sdpa",
    trust_remote_code=True,
).cuda().eval()
processor = AutoProcessor.from_pretrained("henggg/Qwen2.5-VL-3B-CoRL", trust_remote_code=True)

images = [Image.open("view_a.jpg"), Image.open("view_b.jpg")]
messages = [{"role": "user", "content": [
    {"type": "image", "image": images[0]},
    {"type": "image", "image": images[1]},
    {"type": "text", "text": "How many objects are on the table?"},
]}]

text = processor.apply_chat_template(messages, tokenize=False, add_generation_prompt=True)
inputs = processor(text=[text], images=images, return_tensors="pt", padding=True).to("cuda")

gen_ids = model.generate(**inputs, max_new_tokens=2048, do_sample=False)
response = processor.tokenizer.decode(gen_ids[0, inputs["input_ids"].shape[1]:], skip_special_tokens=True)
print(response)

Training from Scratch (8x A100)

git clone https://github.com/hengzzzhou/CoRL.git
cd CoRL
pip install -e .

# SFT -> RL pipeline
torchrun --nproc_per_node=8 scripts/train_production.py \
    --config configs/paper_reproduction.yaml \
    --phase both \
    --output-dir ./outputs

Installation

git clone https://github.com/hengzzzhou/CoRL.git
cd CoRL
pip install -e .

Requirements

  • Python: 3.8+
  • PyTorch: 2.0+ with CUDA (for training)
  • GPU: 8x A100 80GB recommended for paper reproduction
  • Inference works on a single GPU with >= 16GB VRAM

Architecture

CoRL Framework: Stage 1 performs SFT with Chain-of-Thought supervision; Stage 2 applies reinforcement fine-tuning with Cross-View Spatial Reward (CVSR) comprising grounding, overlap, and answer rewards.

Core Components

Module Description
corl/models/qwen_vlm.py Qwen2.5-VL wrapper with native multi-image support
corl/training/sft_trainer.py SFT with chat-template formatting and label masking
corl/training/rl_trainer.py REINFORCE with moving-average baseline
corl/rewards/cvsr_reward.py Multi-component reward (accuracy + grounding + consistency)
corl/data/e2w_dataset.py E2W dataset loader (VERL parquet format)
corl/configs/base_config.py YAML-based configuration system

Dataset

E2W-Bench: Three-level tasks (Counting, Location Reasoning, Grasping) with data from both simulation and real-world environments.

The E2W-Bench dataset contains dual-viewpoint robot images with spatial reasoning tasks:

Task Description Metric
E2W-1: Object Counting "How many X are visible?" Exact match (%)
E2W-2: Spatial Understanding "Which object is second from the right?" Exact match (%)
E2W-3: Grasping "Grasp the X at the specified location" max(0, 1 - dist/d_max) x 100

Each sample includes:

  • 2 JPEG images (robot viewpoints A and B)
  • System + user prompt with <image> placeholders
  • Chain-of-thought ground truth with bounding box annotations

Training

Configuration

Edit configs/paper_reproduction.yaml:

model:
  name: "qwen2.5-vl-3b"
  path: "Qwen/Qwen2.5-VL-3B-Instruct"
  max_seq_len: 1024
  gradient_checkpointing: true

training:
  sft:
    batch_size: 1
    learning_rate: 0.00002
    max_steps: 10000
    grad_accumulation_steps: 8
  rl:
    batch_size: 1
    learning_rate: 0.000001
    max_steps: 5000

data:
  root_path: "./data/CoRL-data"
  sft_max_samples: 40000

Phase 1: SFT

torchrun --nproc_per_node=8 scripts/train_production.py \
    --config configs/paper_reproduction.yaml \
    --phase sft --output-dir ./outputs

Phase 2: RL (after SFT)

torchrun --nproc_per_node=8 scripts/train_production.py \
    --config configs/paper_reproduction.yaml \
    --phase rl --resume ./outputs/sft/sft_final.pt --output-dir ./outputs

Evaluation (Paper Table 1 Metrics)

python scripts/evaluate_paper.py \
    --model-path henggg/Qwen2.5-VL-3B-CoRL \
    --data-path ./data/CoRL-data \
    --split test

Testing

pytest tests/ -v

Project Structure

CoRL/
├── corl/
│   ├── models/        # Qwen2.5-VL integration
│   ├── training/      # Production SFT + RL trainers
│   ├── rewards/       # CVSR reward system
│   ├── data/          # E2W dataset loader
│   ├── configs/       # Configuration system
│   ├── distributed/   # DDP/FSDP utilities
│   └── cli.py         # Command-line interface
├── scripts/
│   ├── train_production.py  # Distributed training entry
│   ├── evaluate_paper.py    # Paper Table 1 evaluation
│   └── inference.py         # Inference with pre-trained models
├── configs/
│   └── paper_reproduction.yaml
└── tests/

Model Zoo

Model HuggingFace E2W-1 E2W-2(S) E2W-2(R) Avg.Reasoning E2W-3(S) Avg.Perception
Qwen2.5-VL-3B + CoRL henggg/Qwen2.5-VL-3B-CoRL 60.8 92.0 84.0 78.9 97.7 97.7
Qwen2.5-VL-7B + CoRL henggg/Qwen2.5-VL-7B-CoRL 66.8 92.8 89.2 82.9 97.8 97.8

E2W-Bench tasks: E2W-1 (Object Counting, exact match %), E2W-2 (Spatial Reasoning, Sim/Real, exact match %), E2W-3 (Grasping, Sim, proximity score)

Citation

@misc{zhou2026egoworldcollaborativespatial,
      title={Ego to World: Collaborative Spatial Reasoning in Embodied Systems via Reinforcement Learning}, 
      author={Heng Zhou and Li Kang and Yiran Qin and Xiufeng Song and Ao Yu and Zilu Zhang and Haoming Song and Kaixin Xu and Yuchen Fan and Dongzhan Zhou and Xiaohong Liu and Ruimao Zhang and Philip Torr and Lei Bai and Zhenfei Yin},
      year={2026},
      eprint={2603.14811},
      archivePrefix={arXiv},
      primaryClass={cs.RO},
      url={https://arxiv.org/abs/2603.14811}, 
}

License

MIT License - see LICENSE.

About

No description, website, or topics provided.

Resources

Stars

Watchers

Forks

Releases

Packages

Contributors

Languages