lerobot/pusht
Viewer β’ Updated β’ 25.7k β’ 14.5k β’ 56
import argparse
import multiprocessing
import sys
from lightning.pytorch.callbacks import LearningRateMonitor, ModelCheckpoint
from lightning.pytorch.loggers import WandbLogger
from physicalai.data import LeRobotDataModule
from physicalai.gyms import PushTGym
from physicalai.policies import Rldx1
from physicalai.train import IterationTimer, Trainer
def parse_args() -> argparse.Namespace:
parser = argparse.ArgumentParser(description="Train RLDX-1 on PushT dataset")
parser.add_argument("--max-epochs", type=int, default=60, help="Number of training epochs")
parser.add_argument(
"--num-workers",
type=int,
default=4,
help="DataLoader workers. Use 0-2 if you see worker/decode failures.",
)
parser.add_argument("--experiment-name", type=str, default=None, help="Name for this wandb experiment run")
return parser.parse_args()
if __name__ == "__main__":
args = parse_args()
# Forked DataLoader workers can deadlock or crash under debugpy; spawn is safer.
multiprocessing.set_start_method("spawn", force=True)
model = Rldx1(
base_model_path="RLWRLD/RLDX-1-PT",
gradient_checkpointing=True,
tune_llm=False,
tune_visual=True,
tune_projector=True,
tune_diffusion_model=True,
color_jitter_params=None,
clip_outliers=False,
tune_top_llm_layers=6,
tune_vlln=False,
video_length=4,
video_stride=1,
n_action_steps=10,
learning_rate=1e-4,
scheduler_decay_lr=1e-5,
max_state_dim=2,
max_action_dim=2,
)
datamodule = LeRobotDataModule(
repo_id="lerobot/pusht",
train_batch_size=8,
data_format="physicalai",
val_gym=PushTGym(),
num_workers=args.num_workers,
)
# Save best checkpoint based on gym reward + keep last 3
best_checkpoint = ModelCheckpoint(
monitor="val/gym/pc_success", # success rate?
mode="max",
save_top_k=2,
filename="rldx1-pusht-{epoch:03d}-{val/gym/pc_success:.2f}",
save_last=True,
verbose=True,
save_weights_only=True,
)
# Log learning rate for debugging schedule issues
lr_monitor = LearningRateMonitor(logging_interval="step")
trainer = Trainer(
max_epochs=args.max_epochs,
accelerator="gpu",
devices=1,
precision="bf16-mixed",
log_every_n_steps=20,
check_val_every_n_epoch=1,
callbacks=[best_checkpoint, lr_monitor, IterationTimer()],
logger=WandbLogger(
project="rldx1-pusht-physical-ai-studio",
name=args.experiment_name,
),
)
trainer.fit(model=model, datamodule=datamodule)
from huggingface_hub import hf_hub_download
from physicalai.policies import Rldx1
from physicalai.gyms import PushTGym
from physicalai.eval.rollout import evaluate_policy
from physicalai.eval.video import VideoRecorder
# Debug switch: set to False to force single-frame inference.
if __name__ == "__main__":
# Download the checkpoint from HuggingFace Hub (cached after first download).
ckpt_path = hf_hub_download(
repo_id="eugene123tw/rldx1_pusht",
filename="pc_success=70.00.ckpt",
)
# Load trained model from checkpoint.
# map_location="cpu" deserializes weights onto CPU first; without it Lightning
# restores tensors onto the checkpoint's saved (cuda) device and can OOM before eval.
model = Rldx1.load_from_checkpoint(ckpt_path, map_location="cpu")
model.eval()
model.cuda()
# Render at the gym/dataset native 96x96. The lerobot/pusht frames the model
# trained on are 96x96; the preprocessor cubically UPSCALES them to 224x224
# (via image_min_area). Rendering at 224 here produces a *sharp* native 224
# image instead of the *blurry* 96->224 upscale training saw -> visual OOD ->
# 0% success. Matching the render to the dataset resolution (96) reproduces
# the exact training input.
env = PushTGym()
# Record videos of all episodes for visualization
recorder = VideoRecorder(
output_dir="./tmp_scripts/pusht_eval_videos",
fps=10,
record_mode="all",
)
# Evaluate over 10 episodes
results = evaluate_policy(
env,
model,
n_episodes=10,
start_seed=0,
video_recorder=recorder,
frame_key="top",
)
recorder.close()
# Print results
agg = results["aggregated"]
print("\n===== Push-T Evaluation Results =====")
print(f"Episodes: {agg['n_episodes']}")
if "pc_success" in agg:
print(f"Success Rate: {agg['pc_success']:.1f}%")
print(f"Num Successes: {agg['num_successes']}")
print(f"Avg Sum Reward: {agg['avg_sum_reward']:.4f}")
print(f"Avg Max Reward: {agg['avg_max_reward']:.4f}")
print(f"Avg Episode Length: {agg['avg_episode_length']:.1f}")
print(f"Avg FPS: {agg['avg_fps']:.1f}")
# Print per-episode breakdown
print("\n----- Per Episode -----")
for ep in results["per_episode"]:
status = "β" if ep.get("success", False) else "β"
print(f" Episode {ep['episode_idx']:3d}: {status} reward={ep['sum_reward']:.4f} steps={ep['episode_length']}")
print(f"\nVideos saved to: ./tmp_scripts/pusht_eval_videos/")
===== Push-T Evaluation Results =====
Episodes: 20
Success Rate: 40.0%
Num Successes: 8
Avg Sum Reward: 96.6334
Avg Max Reward: 0.9024
Avg Episode Length: 248.1
Avg FPS: 59.9
----- Per Episode -----
Episode 0: β reward=132.5254 steps=300
Episode 1: β reward=131.0143 steps=300
Episode 2: β reward=19.8552 steps=104
Episode 3: β reward=158.6325 steps=245
Episode 4: β reward=200.3763 steps=300
Episode 5: β reward=0.0000 steps=300
Episode 6: β reward=127.5933 steps=300
Episode 7: β reward=118.0323 steps=300
Episode 8: β reward=80.0426 steps=300
Episode 9: β reward=8.6257 steps=94
Episode 10: β reward=69.8669 steps=300
Episode 11: β reward=138.0157 steps=300
Episode 12: β reward=50.6851 steps=240
Episode 13: β reward=68.0637 steps=174
Episode 14: β reward=33.0690 steps=85
Episode 15: β reward=134.2909 steps=300
Episode 16: β reward=140.6597 steps=300
Episode 17: β reward=159.8599 steps=300
Episode 18: β reward=88.6393 steps=236
Episode 19: β reward=72.8195 steps=183