| name | vla-rft-world-model-reinforcement-robotics |
| title | VLA-RFT: Reinforcement Fine-Tuning via World Model Simulation |
| version | 0.0.2 |
| engine | skillxiv-v0.0.2-claude-opus-4.6 |
| license | MIT |
| url | https://arxiv.org/abs/2510.00406 |
| keywords | ["VLA","world-models","reinforcement-learning","robotics","efficiency"] |
| description | Fine-tune Vision-Language-Action models using learned world models as simulators, eliminating costly real-world or physics-simulation RL. Train robust robot policies in 400 steps via GRPO with model-generated verified rewards. |
VLA-RFT: Reinforcement Fine-Tuning via World Model Simulation
VLA-RFT addresses the sample inefficiency of robot learning by using a lightweight learned world model (138M parameters) as a simulator to generate verified rewards for policy optimization. This eliminates sim-to-real gaps and real-world cost while enabling rapid fine-tuning.
Core Architecture
- Lightweight world model: 138M-parameter model predicting next visual frames
- Data-driven simulator: Uses model predictions to compute rewards without physics engine
- GRPO-based RL: Advantage-based optimization on simulated trajectories
- Verified rewards: Geometric and kinematic metrics (GIoU, trajectory distance)
- Rapid convergence: 400 training steps sufficient for task learning
Implementation Steps
Train lightweight world model for robotic prediction:
from vla_rft import WorldModel, VLAPolicy, RLTrainer
world_model = WorldModel(
vision_encoder="clip",
prediction_heads=["next_frame", "action_feasibility"],
model_size="lightweight",
prediction_horizon=1
)
world_model.train_on_demonstrations(
data=robot_demonstrations,
learning_rate=1e-4,
num_epochs=10,
reconstruction_loss="mse",
action_feasibility_loss="bce"
)
Execute VLA reinforcement fine-tuning with world model rewards:
vla_policy = VLAPolicy(
base_model="pretrained_vla",
action_space="continuous",
device="cuda"
)
rl_trainer = RLTrainer(
policy=vla_policy,
world_model=world_model,
algorithm="GRPO",
num_rollouts=,
num_training_steps=
)
():
predicted_frames = world_model.predict_frames(
initial_frame=trajectory.initial_image,
actions=trajectory.actions
)
final_predicted_frame = predicted_frames[-]
predicted_object_bbox = extract_bbox(final_predicted_frame)
giou_reward = compute_giou(
predicted=predicted_object_bbox,
target=target.bbox
)
predicted_path = extract_trajectory(predicted_frames)
traj_distance = compute_path_distance(
actual=predicted_path,
optimal=target.optimal_path
)
total_reward = * giou_reward + * ( - traj_distance)
total_reward
step (num_training_steps):
trajectories = rl_trainer.rollout_batch(
num_rollouts=,
task=current_task
)
rewards = []
trajectory trajectories:
reward = compute_reward(trajectory, current_task.target)
rewards.append(reward)
loss = rl_trainer.compute_grpo_loss(
trajectories=trajectories,
rewards=rewards
)
loss.backward()
optimizer.step()
step % == :
()
()