Callback restored
This commit is contained in:
+78
-23
@@ -1,44 +1,99 @@
|
||||
"""
|
||||
ml/callbacks.py - Stable-Baselines3 Custom Callbacks for Logging & Curriculum Advancement
|
||||
Fully compatible with SubprocVecEnv and DummyVecEnv.
|
||||
"""
|
||||
import os
|
||||
|
||||
import numpy as np
|
||||
from stable_baselines3.common.callbacks import BaseCallback
|
||||
|
||||
|
||||
class RewardLoggerCallback(BaseCallback):
|
||||
"""Logs individual reward component averages to TensorBoard."""
|
||||
"""
|
||||
Logs individual reward component averages to TensorBoard and prints
|
||||
the best worker's performance breakdown to the console per iteration.
|
||||
"""
|
||||
|
||||
def __init__(self, verbose: int = 0):
|
||||
def __init__(self, verbose: int = 1):
|
||||
super().__init__(verbose)
|
||||
self.iteration = 0
|
||||
|
||||
def _on_step(self) -> bool:
|
||||
# Pull component averages from the environment vector
|
||||
for env_idx, env in enumerate(self.training_env.envs):
|
||||
if hasattr(env, "get_reward_component_averages"):
|
||||
averages = env.get_reward_component_averages()
|
||||
for key, val in averages.items():
|
||||
self.logger.record(f"reward_components/{key}", val)
|
||||
return True
|
||||
|
||||
def _on_rollout_end(self) -> None:
|
||||
self.iteration += 1
|
||||
if self.training_env is None:
|
||||
return
|
||||
|
||||
try:
|
||||
# Safely query method across all parallel worker processes
|
||||
all_worker_averages = self.training_env.env_method("get_reward_component_averages")
|
||||
except Exception:
|
||||
return
|
||||
|
||||
if not all_worker_averages or len(all_worker_averages) == 0:
|
||||
return
|
||||
|
||||
# 1. Log mean component values across ALL workers to TensorBoard
|
||||
component_keys = all_worker_averages[0].keys()
|
||||
for key in component_keys:
|
||||
mean_val = float(np.mean([w.get(key, 0.0) for w in all_worker_averages]))
|
||||
self.logger.record(f"reward_components/{key}", mean_val)
|
||||
|
||||
# 2. Identify the best performing worker of this iteration
|
||||
worker_totals = [sum(w.values()) for w in all_worker_averages]
|
||||
best_worker_idx = int(np.argmax(worker_totals))
|
||||
best_averages = all_worker_averages[best_worker_idx]
|
||||
best_total = worker_totals[best_worker_idx]
|
||||
|
||||
# 3. Print best worker breakdown to console
|
||||
if self.verbose > 0:
|
||||
print(f"\n" + "=" * 65)
|
||||
print(f" ITERATION {self.iteration} | BEST WORKER (#{best_worker_idx}) REWARD BREAKDOWN")
|
||||
print(f" Total Avg Reward / Step: {best_total:+.4f}")
|
||||
print("-" * 65)
|
||||
for key, val in best_averages.items():
|
||||
print(f" • {key:<26}: {val:+.5f}")
|
||||
print("=" * 65 + "\n")
|
||||
|
||||
|
||||
class CurriculumCallback(BaseCallback):
|
||||
"""Monitors evaluation performance and automatically manages curriculum progression."""
|
||||
"""
|
||||
Monitors training metrics using SB3's native ep_info_buffer and
|
||||
dynamically advances curriculum phases across worker processes.
|
||||
"""
|
||||
|
||||
def __init__(self, eval_freq: int = 10000, verbose: int = 1):
|
||||
def __init__(self, reward_threshold: float = 100.0, verbose: int = 1):
|
||||
super().__init__(verbose)
|
||||
self.eval_freq = eval_freq
|
||||
self.reward_threshold = reward_threshold
|
||||
|
||||
def _on_step(self) -> bool:
|
||||
if self.n_calls % self.eval_freq == 0:
|
||||
for env in self.training_env.envs:
|
||||
if hasattr(env, "get_current_robot_metrics"):
|
||||
metrics = env.get_current_robot_metrics()
|
||||
if metrics:
|
||||
phase = metrics[0].get("phase_name", "UNKNOWN")
|
||||
dist = metrics[0].get("distance_from_start", 0.0)
|
||||
self.logger.record("curriculum/phase_idx", phase)
|
||||
self.logger.record("curriculum/max_distance", dist)
|
||||
if self.verbose > 0:
|
||||
print(f"[CurriculumCallback] Step {self.num_timesteps}: Current Phase = {phase}, Max Dist = {dist:.2f}m")
|
||||
return True
|
||||
|
||||
def _on_rollout_end(self) -> None:
|
||||
if self.training_env is None:
|
||||
return
|
||||
|
||||
# SB3 natively records finished episode stats in self.model.ep_info_buffer
|
||||
if hasattr(self.model, "ep_info_buffer") and len(self.model.ep_info_buffer) > 0:
|
||||
recent_rewards = [ep_info["r"] for ep_info in self.model.ep_info_buffer]
|
||||
mean_reward = float(np.mean(recent_rewards[-50:]))
|
||||
|
||||
try:
|
||||
# Query current phase from worker 0
|
||||
phases = self.training_env.get_attr("curriculum_phase")
|
||||
current_phase = phases[0]
|
||||
|
||||
# Advance curriculum if mean reward exceeds threshold
|
||||
if mean_reward >= self.reward_threshold:
|
||||
if hasattr(current_phase, "next"):
|
||||
next_phase = current_phase.next()
|
||||
if next_phase != current_phase:
|
||||
self.training_env.set_attr("curriculum_phase", next_phase)
|
||||
if self.verbose > 0:
|
||||
print(
|
||||
f"\n[Curriculum] 🚀 Promoted workers to phase: {next_phase.name} "
|
||||
f"(Mean Reward: {mean_reward:.2f})"
|
||||
)
|
||||
except Exception:
|
||||
pass # Keep rollout loop running safely if phase check fails
|
||||
+3
-2
@@ -16,7 +16,7 @@ from stable_baselines3.common.callbacks import CheckpointCallback, EvalCallback
|
||||
|
||||
sys.path.append(str(Path(__file__).resolve().parent.parent))
|
||||
from ml.env import JackBotEnv
|
||||
from ml.callbacks import CurriculumCallback
|
||||
from ml.callbacks import CurriculumCallback, RewardLoggerCallback
|
||||
|
||||
# Silence SB3's UserWarning about SubprocVecEnv vs DummyVecEnv
|
||||
warnings.filterwarnings("ignore", category=UserWarning, module="stable_baselines3")
|
||||
@@ -100,6 +100,7 @@ def main():
|
||||
save_path=args.save_dir,
|
||||
name_prefix=f"jackbot_{ppo_name}",
|
||||
)
|
||||
reward_logger_callback = RewardLoggerCallback(verbose=1)
|
||||
curriculum_callback = CurriculumCallback()
|
||||
|
||||
eval_env = DummyVecEnv([lambda: JackBotEnv(use_gui=False, random_command=True)])
|
||||
@@ -118,7 +119,7 @@ def main():
|
||||
try:
|
||||
model.learn(
|
||||
total_timesteps=args.total_timesteps,
|
||||
callback=[checkpoint_callback, curriculum_callback, eval_callback],
|
||||
callback=[checkpoint_callback, reward_logger_callback, curriculum_callback, eval_callback],
|
||||
progress_bar=True,
|
||||
)
|
||||
final_model_path = os.path.join(args.save_dir, f"jackbot_{ppo_name}_final.zip")
|
||||
|
||||
Reference in New Issue
Block a user