Views
No views yet
1import argparse
2import multiprocessing
3import sys
4
5from lightning.pytorch.callbacks import LearningRateMonitor, ModelCheckpoint
6from lightning.pytorch.loggers import WandbLogger
7
8from physicalai.data import LeRobotDataModule
9from physicalai.gyms import PushTGym
10from physicalai.policies import Rldx1
11from physicalai.train import IterationTimer, Trainer
12
13
14def parse_args() -> argparse.Namespace:
15 parser = argparse.ArgumentParser(description="Train RLDX-1 on PushT dataset")
16 parser.add_argument("--max-epochs", type=int, default=60, help="Number of training epochs")
17 parser.add_argument(
18 "--num-workers",
19 type=int,
20 default=4,
21 help="DataLoader workers. Use 0-2 if you see worker/decode failures.",
22 )
23 parser.add_argument("--experiment-name", type=str, default=None, help="Name for this wandb experiment run")
24 return parser.parse_args()
25
26
27if __name__ == "__main__":
28 args = parse_args()
29
30 # Forked DataLoader workers can deadlock or crash under debugpy; spawn is safer.
31 multiprocessing.set_start_method("spawn", force=True)
32
33 model = Rldx1(
34 base_model_path="RLWRLD/RLDX-1-PT",
35 gradient_checkpointing=True,
36 tune_llm=False,
37 tune_visual=True,
38 tune_projector=True,
39 tune_diffusion_model=True,
40 color_jitter_params=None,
41 clip_outliers=False,
42 tune_top_llm_layers=6,
43 tune_vlln=False,
44 video_length=4,
45 video_stride=1,
46 n_action_steps=10,
47 learning_rate=1e-4,
48 scheduler_decay_lr=1e-5,
49 max_state_dim=2,
50 max_action_dim=2,
51 )
52
53 datamodule = LeRobotDataModule(
54 repo_id="lerobot/pusht",
55 train_batch_size=8,
56 data_format="physicalai",
57 val_gym=PushTGym(),
58 num_workers=args.num_workers,
59 )
60
61 # Save best checkpoint based on gym reward + keep last 3
62 best_checkpoint = ModelCheckpoint(
63 monitor="val/gym/pc_success", # success rate?
64 mode="max",
65 save_top_k=2,
66 filename="rldx1-pusht-{epoch:03d}-{val/gym/pc_success:.2f}",
67 save_last=True,
68 verbose=True,
69 save_weights_only=True,
70 )
71
72 # Log learning rate for debugging schedule issues
73 lr_monitor = LearningRateMonitor(logging_interval="step")
74
75 trainer = Trainer(
76 max_epochs=args.max_epochs,
77 accelerator="gpu",
78 devices=1,
79 precision="bf16-mixed",
80 log_every_n_steps=20,
81 check_val_every_n_epoch=1,
82 callbacks=[best_checkpoint, lr_monitor, IterationTimer()],
83 logger=WandbLogger(
84 project="rldx1-pusht-physical-ai-studio",
85 name=args.experiment_name,
86 ),
87 )
88
89 trainer.fit(model=model, datamodule=datamodule)
901from huggingface_hub import hf_hub_download
2
3from physicalai.policies import Rldx1
4from physicalai.gyms import PushTGym
5from physicalai.eval.rollout import evaluate_policy
6from physicalai.eval.video import VideoRecorder
7
8# Debug switch: set to False to force single-frame inference.
9
10if __name__ == "__main__":
11 # Download the checkpoint from HuggingFace Hub (cached after first download).
12 ckpt_path = hf_hub_download(
13 repo_id="eugene123tw/rldx1_pusht",
14 filename="pc_success=70.00.ckpt",
15 )
16
17 # Load trained model from checkpoint.
18 # map_location="cpu" deserializes weights onto CPU first; without it Lightning
19 # restores tensors onto the checkpoint's saved (cuda) device and can OOM before eval.
20 model = Rldx1.load_from_checkpoint(ckpt_path, map_location="cpu")
21
22 model.eval()
23 model.cuda()
24
25 # Render at the gym/dataset native 96x96. The lerobot/pusht frames the model
26 # trained on are 96x96; the preprocessor cubically UPSCALES them to 224x224
27 # (via image_min_area). Rendering at 224 here produces a *sharp* native 224
28 # image instead of the *blurry* 96->224 upscale training saw -> visual OOD ->
29 # 0% success. Matching the render to the dataset resolution (96) reproduces
30 # the exact training input.
31 env = PushTGym()
32
33 # Record videos of all episodes for visualization
34 recorder = VideoRecorder(
35 output_dir="./tmp_scripts/pusht_eval_videos",
36 fps=10,
37 record_mode="all",
38 )
39
40 # Evaluate over 10 episodes
41 results = evaluate_policy(
42 env,
43 model,
44 n_episodes=10,
45 start_seed=0,
46 video_recorder=recorder,
47 frame_key="top",
48 )
49 recorder.close()
50
51 # Print results
52 agg = results["aggregated"]
53 print("\n===== Push-T Evaluation Results =====")
54 print(f"Episodes: {agg['n_episodes']}")
55 if "pc_success" in agg:
56 print(f"Success Rate: {agg['pc_success']:.1f}%")
57 print(f"Num Successes: {agg['num_successes']}")
58 print(f"Avg Sum Reward: {agg['avg_sum_reward']:.4f}")
59 print(f"Avg Max Reward: {agg['avg_max_reward']:.4f}")
60 print(f"Avg Episode Length: {agg['avg_episode_length']:.1f}")
61 print(f"Avg FPS: {agg['avg_fps']:.1f}")
62
63 # Print per-episode breakdown
64 print("\n----- Per Episode -----")
65 for ep in results["per_episode"]:
66 status = "✓" if ep.get("success", False) else "✗"
67 print(f" Episode {ep['episode_idx']:3d}: {status} reward={ep['sum_reward']:.4f} steps={ep['episode_length']}")
68
69 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