Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
35 changes: 35 additions & 0 deletions scripts/rt1/evaluate_bridge.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,35 @@
import argparse

from simpler_env.evaluation.bridge_tasks import (
widowx_task1_pick_object,
widowx_task2_stack_cube,
widowx_task3_put_object_on_top,
widowx_task4_put_object_in_basket,
)
from simpler_env.policies.rt1.rt1_model import RT1Inference


def parse_args():
parser = argparse.ArgumentParser(description="Run Comprehensive ManiSkill2 Evaluation")
parser.add_argument("--ckpt-path", type=str, required=True, help="Path to the checkpoint to evaluate.")
parser.add_argument("--control-freq", type=int, default=5, help="Set control frequency (default->5)")
return parser.parse_args()


if __name__ == "__main__":
args = parse_args()
ckpt_path = args.ckpt_path

policy = RT1Inference(saved_model_path=ckpt_path, policy_setup="widowx_bridge")

print("Policy initialized. Starting evaluation...")

tasks = [widowx_task1_pick_object, widowx_task2_stack_cube, widowx_task3_put_object_on_top, widowx_task4_put_object_in_basket]

final_scores = []
for task in tasks:
cur_scores = task(env_policy=policy, ckpt_path=args.ckpt_path, control_freq=args.control_freq)
final_scores += cur_scores

print("\nEvaluation finished.")
print(f"Final calculated scores: {final_scores}")
2 changes: 1 addition & 1 deletion simpler_env/evaluation/config.py
Original file line number Diff line number Diff line change
Expand Up @@ -15,7 +15,7 @@ class ManiSkill2Config:
# Required parameters
env_name: str # required=True in argparse
task_name: str
episode_id: int = 0
episode_id: int = None

# Policy settings
policy_model: str = "rt1"
Expand Down
8 changes: 5 additions & 3 deletions simpler_env/evaluation/maniskill2_evaluator.py
Original file line number Diff line number Diff line change
Expand Up @@ -18,6 +18,7 @@
rng = np.random.RandomState(master_seed)
success_threshold_s = 5 # [s]


def run_maniskill2_eval_single_episode(
model,
task_name,
Expand Down Expand Up @@ -305,9 +306,10 @@ def _run_single_evaluation(model, args, control_mode, robot_init_x, robot_init_y
elif args.obj_variation_mode == "episode":
sampled_ids = rng.choice(range(36), size=args.obj_episode_range[1], replace=True)
for idx, obj_episode_id in enumerate(sampled_ids):
if kwargs["episode_id"] is None:
kwargs["episode_id"] = idx
success = run_maniskill2_eval_single_episode(obj_episode_id=obj_episode_id, **kwargs)
kwargs_copy = kwargs.copy()
if kwargs_copy["episode_id"] is None:
kwargs_copy["episode_id"] = idx
success = run_maniskill2_eval_single_episode(obj_episode_id=obj_episode_id, **kwargs_copy)
success_arr.append(success)

elif args.obj_variation_mode == "episode_xy":
Expand Down