Callback restored

This commit is contained in:
2026-08-06 16:15:29 +02:00
parent 346ec9e949
commit 405e3ad5f2
2 changed files with 82 additions and 26 deletions
+79 -24
View File
@@ -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
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