Callback restored
This commit is contained in:
+79
-24
@@ -1,44 +1,99 @@
|
|||||||
"""
|
"""
|
||||||
ml/callbacks.py - Stable-Baselines3 Custom Callbacks for Logging & Curriculum Advancement
|
ml/callbacks.py - Stable-Baselines3 Custom Callbacks for Logging & Curriculum Advancement
|
||||||
|
Fully compatible with SubprocVecEnv and DummyVecEnv.
|
||||||
"""
|
"""
|
||||||
import os
|
|
||||||
import numpy as np
|
import numpy as np
|
||||||
from stable_baselines3.common.callbacks import BaseCallback
|
from stable_baselines3.common.callbacks import BaseCallback
|
||||||
|
|
||||||
|
|
||||||
class RewardLoggerCallback(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)
|
super().__init__(verbose)
|
||||||
|
self.iteration = 0
|
||||||
|
|
||||||
def _on_step(self) -> bool:
|
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
|
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):
|
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)
|
super().__init__(verbose)
|
||||||
self.eval_freq = eval_freq
|
self.reward_threshold = reward_threshold
|
||||||
|
|
||||||
def _on_step(self) -> bool:
|
def _on_step(self) -> bool:
|
||||||
if self.n_calls % self.eval_freq == 0:
|
return True
|
||||||
for env in self.training_env.envs:
|
|
||||||
if hasattr(env, "get_current_robot_metrics"):
|
def _on_rollout_end(self) -> None:
|
||||||
metrics = env.get_current_robot_metrics()
|
if self.training_env is None:
|
||||||
if metrics:
|
return
|
||||||
phase = metrics[0].get("phase_name", "UNKNOWN")
|
|
||||||
dist = metrics[0].get("distance_from_start", 0.0)
|
# SB3 natively records finished episode stats in self.model.ep_info_buffer
|
||||||
self.logger.record("curriculum/phase_idx", phase)
|
if hasattr(self.model, "ep_info_buffer") and len(self.model.ep_info_buffer) > 0:
|
||||||
self.logger.record("curriculum/max_distance", dist)
|
recent_rewards = [ep_info["r"] for ep_info in self.model.ep_info_buffer]
|
||||||
if self.verbose > 0:
|
mean_reward = float(np.mean(recent_rewards[-50:]))
|
||||||
print(f"[CurriculumCallback] Step {self.num_timesteps}: Current Phase = {phase}, Max Dist = {dist:.2f}m")
|
|
||||||
return True
|
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))
|
sys.path.append(str(Path(__file__).resolve().parent.parent))
|
||||||
from ml.env import JackBotEnv
|
from ml.env import JackBotEnv
|
||||||
from ml.callbacks import CurriculumCallback
|
from ml.callbacks import CurriculumCallback, RewardLoggerCallback
|
||||||
|
|
||||||
# Silence SB3's UserWarning about SubprocVecEnv vs DummyVecEnv
|
# Silence SB3's UserWarning about SubprocVecEnv vs DummyVecEnv
|
||||||
warnings.filterwarnings("ignore", category=UserWarning, module="stable_baselines3")
|
warnings.filterwarnings("ignore", category=UserWarning, module="stable_baselines3")
|
||||||
@@ -100,6 +100,7 @@ def main():
|
|||||||
save_path=args.save_dir,
|
save_path=args.save_dir,
|
||||||
name_prefix=f"jackbot_{ppo_name}",
|
name_prefix=f"jackbot_{ppo_name}",
|
||||||
)
|
)
|
||||||
|
reward_logger_callback = RewardLoggerCallback(verbose=1)
|
||||||
curriculum_callback = CurriculumCallback()
|
curriculum_callback = CurriculumCallback()
|
||||||
|
|
||||||
eval_env = DummyVecEnv([lambda: JackBotEnv(use_gui=False, random_command=True)])
|
eval_env = DummyVecEnv([lambda: JackBotEnv(use_gui=False, random_command=True)])
|
||||||
@@ -118,7 +119,7 @@ def main():
|
|||||||
try:
|
try:
|
||||||
model.learn(
|
model.learn(
|
||||||
total_timesteps=args.total_timesteps,
|
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,
|
progress_bar=True,
|
||||||
)
|
)
|
||||||
final_model_path = os.path.join(args.save_dir, f"jackbot_{ppo_name}_final.zip")
|
final_model_path = os.path.join(args.save_dir, f"jackbot_{ppo_name}_final.zip")
|
||||||
|
|||||||
Reference in New Issue
Block a user