aegislm / backend /autotune /optimizer.py
ACA050's picture
Upload 28 files
3e4aa3e verified
Raw
History Blame Contribute Delete
9.17 kB
"""
Optimizer for Autotune Module
Implements policy weight optimization with various algorithms.
Supports gradient descent, Q-learning updates, and weight regularization.
"""
from typing import Dict, Optional, Any, List
from dataclasses import dataclass
from datetime import datetime
import logging
import random
import numpy as np
from .schemas import SystemState, ActionType, PolicyUpdate, TrainingEpisode
from .policy_model import PolicyModel
logger = logging.getLogger(__name__)
@dataclass
class OptimizerConfig:
"""
Configuration for policy optimizer.
"""
learning_rate: float = 0.01
discount_factor: float = 0.95 # Gamma for TD learning
regularization: float = 0.001 # L2 regularization
max_gradient_norm: float = 1.0 # Gradient clipping
update_frequency: int = 1 # Update every N steps
batch_size: int = 32
# Exploration settings
epsilon_start: float = 0.2
epsilon_decay: float = 0.995
epsilon_min: float = 0.01
class Optimizer:
"""
Policy optimizer with Q-learning and gradient descent.
Supports:
- Q-learning updates
- Gradient descent with regularization
- Experience replay
- Target network (for stability)
"""
def __init__(
self,
policy_model: PolicyModel,
config: Optional[OptimizerConfig] = None,
):
"""
Initialize optimizer.
Args:
policy_model: Policy model to optimize
config: Optimizer configuration
"""
self.policy_model = policy_model
self.config = config or OptimizerConfig()
# Experience replay buffer
self.experience_buffer: List[Dict[str, Any]] = []
self.max_buffer_size = 10000
# Target network for stability
self.target_model: Optional[PolicyModel] = None
self.target_update_freq = 100 # Update target every N steps
self.steps_since_update = 0
# Training statistics
self.total_updates = 0
self.update_history: List[Dict[str, float]] = []
def add_experience(
self,
state: SystemState,
action: ActionType,
reward: float,
next_state: Optional[SystemState] = None,
done: bool = False,
) -> None:
"""
Add experience to replay buffer.
Args:
state: Current state
action: Action taken
reward: Reward received
next_state: Next state (if available)
done: Whether episode is done
"""
experience = {
"state": state,
"action": action,
"reward": reward,
"next_state": next_state,
"done": done,
"timestamp": datetime.utcnow(),
}
self.experience_buffer.append(experience)
# Maintain buffer size
if len(self.experience_buffer) > self.max_buffer_size:
self.experience_buffer.pop(0)
def update(
self,
state: SystemState,
action: ActionType,
reward: float,
next_state: Optional[SystemState] = None,
) -> Optional[Dict[str, float]]:
"""
Perform one update step.
Args:
state: Current state
action: Action taken
reward: Reward received
next_state: Next state
Returns:
Update statistics if update performed
"""
# Add to experience buffer
self.add_experience(state, action, reward, next_state)
# Check if we should update
if len(self.experience_buffer) < self.config.batch_size:
return None
# Perform batch update
return self._update_from_batch()
def _update_from_batch(self) -> Dict[str, float]:
"""
Update policy from random batch of experiences.
Returns:
Update statistics
"""
# Sample batch
batch_size = min(
self.config.batch_size,
len(self.experience_buffer)
)
batch = random.sample(self.experience_buffer, batch_size)
total_gradient_norm = 0.0
total_td_error = 0.0
for experience in batch:
state = experience["state"]
action = experience["action"]
reward = experience["reward"]
next_state = experience["next_state"]
# Get weights before update
weights_before = self.policy_model.weights.copy()
# Update weights
gradient_norm = self.policy_model.update_weights(
state=state,
action=action,
reward=reward,
next_state=next_state,
learning_rate=self.config.learning_rate,
)
# Get weights after update
weights_after = self.policy_model.weights.copy()
# Apply regularization
self._apply_regularization()
total_gradient_norm += gradient_norm
# Compute TD error
q_before = float(
self.policy_model.predict_q_values(state)[
list(ActionType).index(action)
]
)
total_td_error += abs(reward - q_before)
# Update target network
self.steps_since_update += 1
if self.steps_since_update >= self.target_update_freq:
self._update_target_network()
self.steps_since_update = 0
# Update statistics
self.total_updates += 1
stats = {
"avg_gradient_norm": total_gradient_norm / batch_size,
"avg_td_error": total_td_error / batch_size,
"buffer_size": len(self.experience_buffer),
}
self.update_history.append(stats)
if len(self.update_history) > 1000:
self.update_history.pop(0)
return stats
def _apply_regularization(self) -> None:
"""
Apply L2 regularization to weights.
"""
# Apply L2 regularization to weights
self.policy_model.weights *= (1.0 - self.config.regularization)
def _update_target_network(self) -> None:
"""
Update target network for stability.
"""
if self.target_model is None:
# Create target network
self.target_model = PolicyModel(
state_dim=self.policy_model.state_dim,
num_actions=self.policy_model.num_actions,
learning_rate=self.policy_model.learning_rate,
epsilon=0.0, # No exploration in target
)
# Copy weights
self.target_model.weights = self.policy_model.weights.copy()
logger.debug("Target network updated")
def get_statistics(self) -> Dict[str, Any]:
"""
Get optimizer statistics.
Returns:
Dictionary of statistics
"""
recent_updates = self.update_history[-100:]
if recent_updates:
avg_gradient = sum(u["avg_gradient_norm"] for u in recent_updates) / len(recent_updates)
avg_td_error = sum(u["avg_td_error"] for u in recent_updates) / len(recent_updates)
else:
avg_gradient = 0.0
avg_td_error = 0.0
return {
"total_updates": self.total_updates,
"buffer_size": len(self.experience_buffer),
"avg_gradient_norm": avg_gradient,
"avg_td_error": avg_td_error,
"epsilon": self.policy_model.strategy.epsilon,
"target_network_ready": self.target_model is not None,
}
def reset(self) -> None:
"""Reset optimizer state."""
self.experience_buffer.clear()
self.total_updates = 0
self.update_history.clear()
self.steps_since_update = 0
# Global optimizer instance
_optimizer: Optional[Optimizer] = None
def get_optimizer(
policy_model: Optional[PolicyModel] = None,
config: Optional[OptimizerConfig] = None,
) -> Optimizer:
"""
Get or create global Optimizer instance.
Args:
policy_model: Policy model to optimize
config: Optimizer configuration
Returns:
Optimizer instance
"""
global _optimizer
if _optimizer is None:
if policy_model is None:
policy_model = PolicyModel()
_optimizer = Optimizer(
policy_model=policy_model,
config=config,
)
return _optimizer
__all__ = [
"Optimizer",
"OptimizerConfig",
"get_optimizer",
]