Files
JackBot/ml/run_eval.py
T
JackM323 c5ca79a354 Machine Learning Trainer
Training environment to make a walk model for the hexapod
generated code that will be checked
2026-07-30 17:23:36 +02:00

40 lines
1.4 KiB
Python

"""Run a trained policy in the PyBullet sim for quick inspection.
Usage:
python ml/run_eval.py --model ml/checkpoints/ppo_joint_command.zip --episodes 3 --gui
"""
import argparse
import sys
from pathlib import Path
ROOT_DIR = Path(__file__).resolve().parent.parent
if str(ROOT_DIR) not in sys.path:
sys.path.insert(0, str(ROOT_DIR))
from ml.evaluate import evaluate
def main():
parser = argparse.ArgumentParser()
parser.add_argument("--model", type=str, required=True, help="Path to the trained model file (.zip)")
parser.add_argument("--episodes", type=int, default=3, help="Number of 'rounds' to run. One episode lasts from reset until the robot falls over or the time limit is reached.")
parser.add_argument("--gui", action="store_true", help="Show the PyBullet GUI during evaluation")
parser.add_argument("--num-robots", type=int, default=1, help="Number of robots in the evaluation environment")
parser.add_argument("--robot-spacing", type=float, default=0.5, help="Spacing between robots in meters")
parser.add_argument("--start-pose", type=str, choices=["init_deg", "init90_deg"], default="init_deg", help="Initial robot pose at reset")
args = parser.parse_args()
evaluate(
model_path=args.model,
episodes=args.episodes,
use_gui=args.gui,
num_robots=args.num_robots,
robot_spacing=args.robot_spacing,
start_pose=args.start_pose,
)
if __name__ == "__main__":
main()