From f684bf44b851b97d12a15dd927d210fe7479396e Mon Sep 17 00:00:00 2001 From: JackM323 Date: Tue, 4 Aug 2026 22:29:30 +0200 Subject: [PATCH] Evaluation Phase Usage --- ml/env.py | 2 +- ml/evaluate.py | 45 ++++++++++++++++++++++++++++++++++----------- ml/run_eval.py | 11 ++++++----- 3 files changed, 41 insertions(+), 17 deletions(-) diff --git a/ml/env.py b/ml/env.py index 2fa971d..14053da 100644 --- a/ml/env.py +++ b/ml/env.py @@ -100,7 +100,7 @@ class JackBotEnv(gym.Env): # Curriculum Initialization via Enum self.curriculum_phase = CurriculumPhase.STAND_ONLY - self.curriculum_episode_limit = 500 # 500 steps limit gives headroom for 400-step requirement + self.curriculum_episode_limit = 300 # 300 steps limit gives headroom for 400-step requirement # Gates required to unlock each target phase self.curriculum_stage_requirements = { diff --git a/ml/evaluate.py b/ml/evaluate.py index 2303859..c3e413f 100644 --- a/ml/evaluate.py +++ b/ml/evaluate.py @@ -1,5 +1,5 @@ """ -ml/evaluate.py - Evaluation routine for trained JackBot PPO policies. +ml/evaluate.py - Curriculum-aware evaluation routine for trained JackBot PPO policies. """ import argparse import json @@ -15,7 +15,7 @@ def evaluate( use_gui: bool = True, robot_spacing: float = 0.5, start_pose: str = "init_deg", - random_command: bool = False, + random_command: bool = True, # Default to True so curriculum commands are sampled save_json: Optional[str] = None, ) -> Dict[str, Any]: try: @@ -31,7 +31,7 @@ def evaluate( # Force device="cpu" to prevent AMD ROCm/hipBLASLt matrix multiplication crashes model = PPO.load(model_path, device="cpu") - # Initialize standard environment (returns single array of shape (352,)) + # Initialize environment with random commands enabled for curriculum progression env = JackBotEnv( use_gui=use_gui, random_command=random_command, @@ -41,45 +41,68 @@ def evaluate( episode_rewards: List[float] = [] episode_lengths: List[int] = [] + episode_phases: List[str] = [] for ep in range(episodes): obs, _ = env.reset() - terminated = False + done = False total_reward = 0.0 steps = 0 - print(f"\n--- Starting Evaluation Episode {ep + 1}/{episodes} ---") + initial_phase = env.curriculum_phase.name + print(f"\n--- Starting Evaluation Episode {ep + 1}/{episodes} [Phase: {initial_phase}] ---") - while not terminated: + while not done: action, _ = model.predict(obs, deterministic=True) + + # Store phase before step to detect live phase transitions + prev_phase = env.curriculum_phase + obs, reward, terminated, truncated, _ = env.step(action) + # Done on either physical failure (terminated) or phase step limit (truncated) + done = terminated or truncated total_reward += float(reward) steps += 1 + # Log live phase transition if unlocked during this step + if env.curriculum_phase != prev_phase: + print(f" └─ [Eval Milestone] Curriculum advanced to {env.curriculum_phase.name} at episode step {steps}!") + if use_gui: time.sleep(1.0 / 240.0) + status_str = "FAILED (Terminated)" if terminated else "COMPLETED (Truncated)" + final_phase = env.curriculum_phase.name + episode_rewards.append(total_reward) episode_lengths.append(steps) - print(f"Episode {ep + 1} Finished: Total Reward = {total_reward:.2f} | Steps = {steps}") + episode_phases.append(final_phase) + + print( + f"Episode {ep + 1} Finished [{status_str}]: " + f"Phase = {final_phase} | Total Reward = {total_reward:.2f} | Steps = {steps}" + ) env.close() metrics = { "model_path": str(model_path), "episodes_evaluated": episodes, + "final_curriculum_phase": env.curriculum_phase.name, "mean_reward": float(np.mean(episode_rewards)), "std_reward": float(np.std(episode_rewards)), "mean_episode_length": float(np.mean(episode_lengths)), "raw_rewards": episode_rewards, + "episode_phases": episode_phases, } - print("\n" + "=" * 50) + print("\n" + "=" * 60) print(f"EVALUATION COMPLETE ({episodes} Episodes)") - print(f"Mean Reward: {metrics['mean_reward']:.2f} ± {metrics['std_reward']:.2f}") - print(f"Mean Episode Length: {metrics['mean_episode_length']:.1f} steps") - print("=" * 50) + print(f"Final Reached Phase: {metrics['final_curriculum_phase']}") + print(f"Mean Reward: {metrics['mean_reward']:.2f} ± {metrics['std_reward']:.2f}") + print(f"Mean Episode Length: {metrics['mean_episode_length']:.1f} steps") + print("=" * 60) if save_json: out_path = Path(save_json) diff --git a/ml/run_eval.py b/ml/run_eval.py index b02ac82..5bd9f74 100644 --- a/ml/run_eval.py +++ b/ml/run_eval.py @@ -1,7 +1,7 @@ -"""Run a trained policy in the PyBullet sim for quick inspection. +"""Run a trained policy in the PyBullet sim with full curriculum progression. Usage: - python ml/run_eval.py --model ml/checkpoints/ppo_joint_command.zip --episodes 3 --gui + python ml/run_eval.py --model ml/checkpoints/ppo_joint_command.zip --episodes 5 --gui """ import argparse @@ -16,13 +16,14 @@ from ml.evaluate import evaluate def main(): - parser = argparse.ArgumentParser(description="JackBot Policy Evaluator Wrapper") + parser = argparse.ArgumentParser(description="JackBot Curriculum Policy Evaluator Wrapper") 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 evaluation episodes to run.") + parser.add_argument("--episodes", type=int, default=5, help="Number of evaluation episodes to run.") parser.add_argument("--gui", action="store_true", help="Show the PyBullet GUI during evaluation") 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") - parser.add_argument("--random-command", action="store_true", help="Randomize command samples during evaluation") + parser.add_argument("--no-random-command", dest="random_command", action="store_false", help="Disable command sampling and lock to zero commands") + parser.set_defaults(random_command=True) parser.add_argument("--save-metrics", type=str, default=None, help="Optional path to save JSON metrics report") args = parser.parse_args()