| |
|
|
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| import logging |
|
|
| from lerobot.cameras import opencv |
| from lerobot.configs import parser |
| from lerobot.datasets import LeRobotDataset |
| from lerobot.policies import make_policy |
| from lerobot.robots import ( |
| RobotConfig, |
| make_robot_from_config, |
| so_follower, |
| ) |
| from lerobot.teleoperators import ( |
| gamepad, |
| so_leader, |
| ) |
|
|
| from .gym_manipulator import make_robot_env |
| from .train_rl import TrainRLServerPipelineConfig |
|
|
| logging.basicConfig(level=logging.INFO) |
|
|
|
|
| def eval_policy(env, policy, n_episodes): |
| sum_reward_episode = [] |
| for _ in range(n_episodes): |
| obs, _ = env.reset() |
| episode_reward = 0.0 |
| while True: |
| action = policy.select_action(obs) |
| obs, reward, terminated, truncated, _ = env.step(action) |
| episode_reward += reward |
| if terminated or truncated: |
| break |
| sum_reward_episode.append(episode_reward) |
|
|
| logging.info(f"Success after 20 steps {sum_reward_episode}") |
| logging.info(f"success rate {sum(sum_reward_episode) / len(sum_reward_episode)}") |
|
|
|
|
| @parser.wrap() |
| def main(cfg: TrainRLServerPipelineConfig): |
| env_cfg = cfg.env |
| env = make_robot_env(env_cfg) |
| dataset_cfg = cfg.dataset |
| dataset = LeRobotDataset(repo_id=dataset_cfg.repo_id) |
| dataset_meta = dataset.meta |
|
|
| policy = make_policy( |
| cfg=cfg.policy, |
| |
| ds_meta=dataset_meta, |
| ) |
| policy = policy.from_pretrained(env_cfg.pretrained_policy_name_or_path) |
| policy.eval() |
|
|
| eval_policy(env, policy=policy, n_episodes=10) |
|
|
|
|
| if __name__ == "__main__": |
| main() |
|
|