aegislm / backend /autotune /training_loop.py
ACA050's picture
Upload 28 files
3e4aa3e verified
Raw
History Blame Contribute Delete
12.7 kB
"""
Training Loop for Autotune Module
Implements periodic training loop for policy optimization.
Manages training episodes and state-action-reward history.
"""
from typing import Dict, List, Optional, Any, Callable
from dataclasses import dataclass
from datetime import datetime, timedelta
import logging
import threading
import time
from .schemas import (
SystemState,
ActionType,
SchedulerAction,
Reward,
TrainingEpisode,
AdaptiveConfig,
PolicyUpdate,
)
from .state_encoder import StateEncoder
from .action_space import ActionSpace
from .reward_engine import RewardEngine
from .policy_model import PolicyModel
from .optimizer import Optimizer, OptimizerConfig
logger = logging.getLogger(__name__)
@dataclass
class TrainingConfig:
"""
Configuration for training loop.
"""
# Training frequency
training_interval_seconds: int = 300 # 5 minutes
episode_length_steps: int = 100
# Batch settings
batch_size: int = 32
buffer_size: int = 10000
# Exploration
epsilon_start: float = 0.2
epsilon_decay: float = 0.995
epsilon_min: float = 0.01
# Learning
learning_rate: float = 0.01
discount_factor: float = 0.95
# Safety
enable_rollback: bool = True
rollback_threshold: float = 0.1
# Logging
log_interval_seconds: int = 60
save_interval_episodes: int = 100
class TrainingLoop:
"""
Manages the RL training loop for policy optimization.
Runs periodic training updates:
1. Collect state-action-reward tuples
2. Apply scheduling policy
3. Observe reward
4. Update policy weights
5. Log performance
"""
def __init__(
self,
state_encoder: StateEncoder,
action_space: ActionSpace,
reward_engine: RewardEngine,
policy_model: PolicyModel,
optimizer: Optimizer,
config: Optional[TrainingConfig] = None,
):
"""
Initialize training loop.
Args:
state_encoder: State encoder instance
action_space: Action space instance
reward_engine: Reward engine instance
policy_model: Policy model instance
optimizer: Optimizer instance
config: Training configuration
"""
self.state_encoder = state_encoder
self.action_space = action_space
self.reward_engine = reward_engine
self.policy_model = policy_model
self.optimizer = optimizer
self.config = config or TrainingConfig()
# Training state
self.current_episode: Optional[TrainingEpisode] = None
self.episode_count = 0
self.step_count = 0
# Historical data
self.episodes: List[TrainingEpisode] = []
self.policy_updates: List[PolicyUpdate] = []
# State tracking
self.last_state: Optional[SystemState] = None
self.last_action: Optional[SchedulerAction] = None
self.last_reward: Optional[float] = None
# Control flags
self._running = False
self._paused = False
self._stop_event = threading.Event()
# Callbacks
self.on_episode_complete: Optional[Callable] = None
self.on_training_update: Optional[Callable] = None
def start(self) -> None:
"""Start the training loop."""
if self._running:
logger.warning("Training loop already running")
return
self._running = True
self._stop_event.clear()
# Start training thread
self._training_thread = threading.Thread(
target=self._training_loop,
daemon=True
)
self._training_thread.start()
logger.info("Training loop started")
def stop(self) -> None:
"""Stop the training loop."""
if not self._running:
return
self._stop_event.set()
self._running = False
if hasattr(self, '_training_thread'):
self._training_thread.join(timeout=5.0)
logger.info("Training loop stopped")
def pause(self) -> None:
"""Pause the training loop."""
self._paused = True
logger.info("Training loop paused")
def resume(self) -> None:
"""Resume the training loop."""
self._paused = False
logger.info("Training loop resumed")
def record_step(
self,
state: SystemState,
action: SchedulerAction,
reward: float,
next_state: Optional[SystemState] = None,
) -> None:
"""
Record a training step.
Args:
state: Current state
action: Action taken
reward: Reward received
next_state: Next state (if available)
"""
# Start new episode if needed
if self.current_episode is None:
self._start_new_episode()
# Add step to current episode
self.current_episode.add_step(state, action, reward)
self.step_count += 1
# Update optimizer
self.optimizer.update(
state=state,
action=action.action_type,
reward=reward,
next_state=next_state,
)
# Check if episode is complete
if self.current_episode.episode_length >= self.config.episode_length_steps:
self._complete_episode()
# Store last transition
self.last_state = state
self.last_action = action
self.last_reward = reward
def get_current_state(self) -> SystemState:
"""
Get current system state for decision making.
Returns:
Current system state
"""
return self.last_state or self.state_encoder.get_default_state()
def _training_loop(self) -> None:
"""Main training loop (runs in separate thread)."""
last_log_time = time.time()
while not self._stop_event.is_set():
try:
if self._paused:
time.sleep(1.0)
continue
current_time = time.time()
# Log progress periodically
if current_time - last_log_time >= self.config.log_interval_seconds:
self._log_progress()
last_log_time = current_time
# Sleep briefly to avoid busy waiting
time.sleep(1.0)
except Exception as e:
logger.error(f"Error in training loop: {e}", exc_info=True)
logger.info("Training loop thread exited")
def _start_new_episode(self) -> None:
"""Start a new training episode."""
self.episode_count += 1
self.current_episode = TrainingEpisode(
episode_id=self.episode_count,
start_time=datetime.utcnow(),
)
logger.debug(f"Started episode {self.episode_count}")
def _complete_episode(self) -> None:
"""Complete the current episode."""
if self.current_episode is None:
return
# Finalize episode
self.current_episode.finalize()
# Store episode
self.episodes.append(self.current_episode)
# Maintain max episodes
if len(self.episodes) > 1000:
self.episodes.pop(0)
# Log episode summary
logger.info(
f"Episode {self.episode_count} complete: "
f"reward={self.current_episode.total_reward:.4f}, "
f"length={self.current_episode.episode_length}, "
f"unique_actions={self.current_episode.unique_actions}"
)
# Trigger callback
if self.on_episode_complete:
self.on_episode_complete(self.current_episode)
# Start new episode
self._start_new_episode()
def _log_progress(self) -> None:
"""Log training progress."""
stats = self.get_statistics()
logger.info(
f"Training progress: "
f"episodes={stats['episode_count']}, "
f"steps={stats['step_count']}, "
f"epsilon={stats['epsilon']:.4f}, "
f"avg_reward={stats['avg_reward']:.4f}"
)
def get_statistics(self) -> Dict[str, Any]:
"""
Get training statistics.
Returns:
Dictionary of statistics
"""
# Calculate recent average reward
recent_episodes = self.episodes[-10:]
if recent_episodes:
avg_reward = sum(e.avg_reward for e in recent_episodes) / len(recent_episodes)
else:
avg_reward = 0.0
return {
"episode_count": self.episode_count,
"step_count": self.step_count,
"epsilon": self.policy_model.strategy.epsilon,
"avg_reward": avg_reward,
"buffer_size": len(self.optimizer.experience_buffer),
"total_updates": self.optimizer.total_updates,
"is_running": self._running,
"is_paused": self._paused,
}
def get_episode_history(
self,
count: int = 100,
) -> List[TrainingEpisode]:
"""
Get recent episode history.
Args:
count: Number of recent episodes to return
Returns:
List of recent episodes
"""
return self.episodes[-count:]
def decay_epsilon(self) -> None:
"""Decay exploration rate."""
self.policy_model.strategy.decay_epsilon()
logger.debug(
f"Epsilon decayed to {self.policy_model.strategy.epsilon:.4f}"
)
def reset(self) -> None:
"""Reset training state."""
self.stop()
self.current_episode = None
self.episode_count = 0
self.step_count = 0
self.episodes.clear()
self.policy_updates.clear()
self.last_state = None
self.last_action = None
self.last_reward = None
self.optimizer.reset()
logger.info("Training loop reset")
# Global training loop instance
_training_loop: Optional[TrainingLoop] = None
def get_training_loop(
state_encoder: Optional[StateEncoder] = None,
action_space: Optional[ActionSpace] = None,
reward_engine: Optional[RewardEngine] = None,
policy_model: Optional[PolicyModel] = None,
optimizer: Optional[Optimizer] = None,
config: Optional[TrainingConfig] = None,
) -> TrainingLoop:
"""
Get or create global TrainingLoop instance.
Args:
state_encoder: State encoder instance
action_space: Action space instance
reward_engine: Reward engine instance
policy_model: Policy model instance
optimizer: Optimizer instance
config: Training configuration
Returns:
TrainingLoop instance
"""
global _training_loop
if _training_loop is None:
# Create components if not provided
if state_encoder is None:
state_encoder = StateEncoder()
if action_space is None:
action_space = ActionSpace()
if reward_engine is None:
reward_engine = RewardEngine()
if policy_model is None:
policy_model = PolicyModel()
if optimizer is None:
opt_config = OptimizerConfig() if config is None else OptimizerConfig(
learning_rate=config.learning_rate,
discount_factor=config.discount_factor,
)
optimizer = Optimizer(policy_model=policy_model, config=opt_config)
_training_loop = TrainingLoop(
state_encoder=state_encoder,
action_space=action_space,
reward_engine=reward_engine,
policy_model=policy_model,
optimizer=optimizer,
config=config,
)
return _training_loop
__all__ = [
"TrainingLoop",
"TrainingConfig",
"get_training_loop",
]