Files
foxhunt/services/ml_training_service/src/batch_tuning_manager.rs
jgrusewski 74980c47b0 fix(ml, ml_training_service): add error logging to swallowed RwLock poison and channel send failures
- ensemble/model.rs: replace silent .read().ok() with map_err+tracing::error for both models and signals RwLock in health_check
- batch_tuning_manager.rs: replace let _ = stop_tuning_job() with if let Err + tracing::warn
- orchestrator.rs: replace let _ = broadcaster.send() with if let Err + tracing::warn
- tuning_manager.rs: replace let _ = progress_tx.send() with if let Err + tracing::warn
- trial_executor.rs: replace let _ = result_tx.send() with if let Err + tracing::warn
- Use {:?} for SendError types whose inner value does not implement Display

Co-Authored-By: Claude Opus 4.6 <noreply@anthropic.com>
2026-02-21 23:54:25 +01:00

756 lines
28 KiB
Rust

//! Batch Tuning Manager - Automated Multi-Model Hyperparameter Optimization
//!
//! This module orchestrates batch tuning jobs across multiple ML models with:
//! - Automatic dependency resolution (e.g., TFT depends on MAMBA_2)
//! - Sequential execution with smart ordering
//! - Automatic YAML export to ml/config/best_hyperparameters.yaml
//! - Consolidated reporting for model comparison
//!
//! # Architecture
//!
//! ```text
//! BatchTuningManager
//! │
//! ├─ Dependency Resolver (models → execution order)
//! │
//! ├─ Sequential Executor (one model at a time)
//! │ │
//! │ └─ TuningManager (per-model Optuna subprocess)
//! │
//! ├─ YAML Exporter (best_hyperparameters.yaml)
//! │
//! └─ Consolidated Reporter (comparison & recommendations)
//! ```
use anyhow::{anyhow, Context, Result};
use chrono::{DateTime, Utc};
use std::collections::HashMap;
use std::path::Path;
use std::sync::Arc;
use tokio::fs;
use tokio::sync::RwLock;
use tracing::{debug, error, info, warn};
use uuid::Uuid;
use crate::tuning_manager::{TuningJobStatus, TuningManagerTrait};
/// Status of a batch tuning job
#[derive(Debug, Clone, PartialEq)]
pub enum BatchJobStatus {
Pending,
Running,
Completed,
PartiallyCompleted,
Failed,
Stopped,
}
/// Individual model tuning result
#[derive(Debug, Clone)]
pub struct ModelTuningResult {
pub model_type: String,
pub job_id: Uuid,
pub status: TuningJobStatus,
pub best_params: HashMap<String, f32>,
pub best_metrics: HashMap<String, f32>,
pub trials_completed: u32,
pub started_at: DateTime<Utc>,
pub completed_at: Option<DateTime<Utc>>,
pub error_message: Option<String>,
}
/// Batch tuning job metadata
#[derive(Debug, Clone)]
pub struct BatchTuningJob {
pub batch_id: Uuid,
pub status: BatchJobStatus,
pub models: Vec<String>,
pub execution_order: Vec<String>,
pub current_model_index: usize,
pub results: Vec<ModelTuningResult>,
pub trials_per_model: u32,
pub config_path: String,
pub auto_export_yaml: bool,
pub yaml_export_path: String,
pub started_at: DateTime<Utc>,
pub updated_at: DateTime<Utc>,
pub completed_at: Option<DateTime<Utc>>,
}
/// Model dependency rules
/// Format: (dependent_model, required_model)
/// Example: ("TFT", "MAMBA_2") means TFT must run after MAMBA_2
const MODEL_DEPENDENCIES: &[(&str, &str)] = &[
("TFT", "MAMBA_2"), // TFT uses MAMBA_2 features/embeddings
];
pub struct BatchTuningManager {
/// Active batch jobs indexed by batch_id
jobs: Arc<RwLock<HashMap<Uuid, BatchTuningJob>>>,
/// Reference to single-model tuning manager
tuning_manager: Arc<dyn TuningManagerTrait>,
/// Working directory for batch artifacts
working_dir: String,
}
impl BatchTuningManager {
/// Create a new batch tuning manager
pub fn new(tuning_manager: Arc<dyn TuningManagerTrait>, working_dir: String) -> Self {
Self {
jobs: Arc::new(RwLock::new(HashMap::new())),
tuning_manager,
working_dir,
}
}
/// Start a batch tuning job for multiple models
///
/// # Arguments
/// * `models` - List of model types to tune (e.g., ["DQN", "PPO", "MAMBA_2", "TFT"])
/// * `trials_per_model` - Number of Optuna trials for each model
/// * `config_path` - Path to tuning configuration YAML
/// * `data_source_path` - Optional data source path (defaults to test data)
/// * `auto_export_yaml` - Automatically export best params to YAML
/// * `yaml_export_path` - Custom YAML export path (default: ml/config/best_hyperparameters.yaml)
///
/// # Returns
/// Batch job ID (UUID)
pub async fn start_batch_tuning(
&self,
models: Vec<String>,
trials_per_model: u32,
config_path: String,
data_source_path: Option<String>,
auto_export_yaml: bool,
yaml_export_path: Option<String>,
) -> Result<Uuid> {
info!("Starting batch tuning job for models: {:?}", models);
// Validate models
for model in &models {
if !self.is_valid_model(model) {
return Err(anyhow!("Invalid model type: {}", model));
}
}
// Resolve dependencies and determine execution order
let execution_order = self.resolve_model_dependencies(&models);
info!(
"Execution order after dependency resolution: {:?}",
execution_order
);
// Create batch job metadata
let batch_id = Uuid::new_v4();
let yaml_path = yaml_export_path.unwrap_or_else(|| {
let cwd = std::env::current_dir().unwrap_or_else(|_| std::path::PathBuf::from("."));
format!("{}/ml/config/best_hyperparameters.yaml", cwd.display())
});
let job = BatchTuningJob {
batch_id,
status: BatchJobStatus::Pending,
models: models.clone(),
execution_order: execution_order.clone(),
current_model_index: 0,
results: Vec::new(),
trials_per_model,
config_path: config_path.clone(),
auto_export_yaml,
yaml_export_path: yaml_path,
started_at: Utc::now(),
updated_at: Utc::now(),
completed_at: None,
};
// Store job metadata
{
let mut jobs = self.jobs.write().await;
jobs.insert(batch_id, job.clone());
}
// Spawn background task to execute models sequentially
let tuning_manager = Arc::clone(&self.tuning_manager);
let jobs = Arc::clone(&self.jobs);
let working_dir = self.working_dir.clone();
tokio::spawn(async move {
Self::execute_batch_sequentially(
batch_id,
execution_order,
trials_per_model,
config_path,
data_source_path,
tuning_manager,
jobs,
working_dir,
)
.await;
});
// Update status to Running
{
let mut jobs = self.jobs.write().await;
if let Some(job) = jobs.get_mut(&batch_id) {
job.status = BatchJobStatus::Running;
job.updated_at = Utc::now();
}
}
Ok(batch_id)
}
/// Get batch job status
pub async fn get_batch_status(&self, batch_id: Uuid) -> Result<BatchTuningJob> {
let jobs = self.jobs.read().await;
jobs.get(&batch_id)
.cloned()
.ok_or_else(|| anyhow!("Batch job {} not found", batch_id))
}
/// Stop a running batch job
pub async fn stop_batch_job(&self, batch_id: Uuid, reason: String) -> Result<()> {
info!("Stopping batch job {}: {}", batch_id, reason);
let mut jobs = self.jobs.write().await;
if let Some(job) = jobs.get_mut(&batch_id) {
// Stop current tuning job if running
if job.current_model_index < job.execution_order.len() {
if let Some(result) = job.results.last() {
if let Err(e) = self
.tuning_manager
.stop_tuning_job(result.job_id, reason.clone())
.await
{
tracing::warn!("Failed to stop tuning job during batch stop: {}", e);
}
}
}
job.status = BatchJobStatus::Stopped;
job.updated_at = Utc::now();
job.completed_at = Some(Utc::now());
}
Ok(())
}
/// Resolve model dependencies and determine execution order
///
/// # Logic
/// 1. Build dependency graph from MODEL_DEPENDENCIES
/// 2. Topological sort to determine execution order
/// 3. Models without dependencies can run in any order
///
/// # Example
/// Input: ["TFT", "DQN", "MAMBA_2", "PPO"]
/// Output: ["DQN", "PPO", "MAMBA_2", "TFT"] (DQN/PPO first, MAMBA_2 before TFT)
pub fn resolve_model_dependencies(&self, models: &[String]) -> Vec<String> {
let models_set: std::collections::HashSet<_> = models.iter().map(|s| s.as_str()).collect();
let mut graph: HashMap<&str, Vec<&str>> = HashMap::new();
let mut in_degree: HashMap<&str, usize> = HashMap::new();
// Initialize in-degree for all models
for model in models {
in_degree.insert(model.as_str(), 0);
graph.entry(model.as_str()).or_default();
}
// Build dependency graph
for &(dependent, required) in MODEL_DEPENDENCIES {
if models_set.contains(dependent) && models_set.contains(required) {
graph
.entry(required)
.or_default()
.push(dependent);
*in_degree.entry(dependent).or_insert(0) += 1;
}
}
// Topological sort (Kahn's algorithm)
let mut queue: Vec<&str> = in_degree
.iter()
.filter(|(_, &deg)| deg == 0)
.map(|(&model, _)| model)
.collect();
let mut result = Vec::new();
while let Some(model) = queue.pop() {
result.push(model.to_string());
if let Some(dependents) = graph.get(model) {
for &dependent in dependents {
if let Some(deg) = in_degree.get_mut(dependent) {
*deg -= 1;
if *deg == 0 {
queue.push(dependent);
}
}
}
}
}
// Check for cycles (should not happen with our fixed dependency rules)
if result.len() != models.len() {
warn!("Cycle detected in model dependencies, using original order");
return models.to_vec();
}
result
}
/// Execute models sequentially (private background task)
async fn execute_batch_sequentially(
batch_id: Uuid,
execution_order: Vec<String>,
trials_per_model: u32,
config_path: String,
_data_source_path: Option<String>,
tuning_manager: Arc<dyn TuningManagerTrait>,
jobs: Arc<RwLock<HashMap<Uuid, BatchTuningJob>>>,
_working_dir: String,
) {
info!(
"Executing batch {} with {} models sequentially",
batch_id,
execution_order.len()
);
for (index, model_type) in execution_order.iter().enumerate() {
info!(
"Starting tuning for model {} ({}/{})",
model_type,
index + 1,
execution_order.len()
);
// Update current model index
{
let mut jobs_guard = jobs.write().await;
if let Some(job) = jobs_guard.get_mut(&batch_id) {
job.current_model_index = index;
job.updated_at = Utc::now();
}
}
let model_start_time = Utc::now();
// Start tuning job for this model
let tuning_result = tuning_manager
.start_tuning_job(
model_type.clone(),
trials_per_model,
config_path.clone(),
format!("Batch {} - Model {}", batch_id, model_type),
HashMap::new(),
)
.await;
match tuning_result {
Ok(job_id) => {
info!("Tuning job started for {}: {}", model_type, job_id);
// Poll for completion
let final_status =
Self::poll_tuning_completion(&tuning_manager, job_id, trials_per_model)
.await;
let model_result = ModelTuningResult {
model_type: model_type.clone(),
job_id,
status: final_status.status.clone(),
best_params: final_status.best_params,
best_metrics: final_status.best_metrics,
trials_completed: final_status.current_trial,
started_at: model_start_time,
completed_at: Some(Utc::now()),
error_message: final_status.error_message,
};
// Store result
let mut jobs_guard = jobs.write().await;
if let Some(job) = jobs_guard.get_mut(&batch_id) {
job.results.push(model_result);
job.updated_at = Utc::now();
}
},
Err(e) => {
error!("Failed to start tuning job for {}: {}", model_type, e);
let model_result = ModelTuningResult {
model_type: model_type.clone(),
job_id: Uuid::nil(),
status: TuningJobStatus::Failed,
best_params: HashMap::new(),
best_metrics: HashMap::new(),
trials_completed: 0,
started_at: model_start_time,
completed_at: Some(Utc::now()),
error_message: Some(format!("{}", e)),
};
let mut jobs_guard = jobs.write().await;
if let Some(job) = jobs_guard.get_mut(&batch_id) {
job.results.push(model_result);
job.updated_at = Utc::now();
}
},
}
}
// Batch complete - determine final status
let mut jobs_guard = jobs.write().await;
if let Some(job) = jobs_guard.get_mut(&batch_id) {
let total_models = job.results.len();
let successful = job
.results
.iter()
.filter(|r| r.status == TuningJobStatus::Completed)
.count();
let _failed = job
.results
.iter()
.filter(|r| r.status == TuningJobStatus::Failed)
.count();
job.status = if successful == total_models {
BatchJobStatus::Completed
} else if successful > 0 {
BatchJobStatus::PartiallyCompleted
} else {
BatchJobStatus::Failed
};
job.completed_at = Some(Utc::now());
job.updated_at = Utc::now();
info!(
"Batch {} completed: {}/{} models successful, status: {:?}",
batch_id, successful, total_models, job.status
);
// Auto-export YAML if enabled
if job.auto_export_yaml && successful > 0 {
if let Err(e) = Self::export_yaml_internal(job).await {
error!("Failed to auto-export YAML for batch {}: {}", batch_id, e);
}
}
}
}
/// Poll tuning job until completion (internal helper)
async fn poll_tuning_completion(
tuning_manager: &Arc<dyn TuningManagerTrait>,
job_id: Uuid,
expected_trials: u32,
) -> crate::tuning_manager::TuningJob {
loop {
tokio::time::sleep(tokio::time::Duration::from_secs(10)).await;
match tuning_manager.get_tuning_job_status(job_id).await {
Ok(status) => {
debug!(
"Tuning job {} status: {:?}, trials: {}/{}",
job_id, status.status, status.current_trial, expected_trials
);
match status.status {
TuningJobStatus::Completed
| TuningJobStatus::Failed
| TuningJobStatus::Stopped => {
return status;
},
_ => continue,
}
},
Err(e) => {
error!("Failed to get tuning job status for {}: {}", job_id, e);
tokio::time::sleep(tokio::time::Duration::from_secs(5)).await;
},
}
}
}
/// Export best hyperparameters to YAML file
pub async fn export_best_hyperparameters(
&self,
batch_id: Uuid,
_output_path: &str,
) -> Result<()> {
let jobs = self.jobs.read().await;
let job = jobs
.get(&batch_id)
.ok_or_else(|| anyhow!("Batch job {} not found", batch_id))?;
Self::export_yaml_internal(job).await
}
/// Internal YAML export implementation
async fn export_yaml_internal(job: &BatchTuningJob) -> Result<()> {
let output_path = &job.yaml_export_path;
info!("Exporting best hyperparameters to: {}", output_path);
// Create parent directory if it doesn't exist
if let Some(parent) = Path::new(output_path).parent() {
fs::create_dir_all(parent)
.await
.context("Failed to create parent directory")?;
}
// Build YAML content
let mut yaml_content = String::from("# Best Hyperparameters from Batch Tuning\n");
yaml_content.push_str(&format!("# Batch ID: {}\n", job.batch_id));
yaml_content.push_str(&format!("# Generated: {}\n\n", Utc::now().to_rfc3339()));
yaml_content.push_str("models:\n");
for result in &job.results {
if result.status == TuningJobStatus::Completed && !result.best_params.is_empty() {
yaml_content.push_str(&format!(" {}:\n", result.model_type));
// Best hyperparameters
yaml_content.push_str(" hyperparameters:\n");
for (param, value) in &result.best_params {
yaml_content.push_str(&format!(" {}: {}\n", param, value));
}
// Best metrics
yaml_content.push_str(" metrics:\n");
for (metric, value) in &result.best_metrics {
yaml_content.push_str(&format!(" {}: {:.6}\n", metric, value));
}
yaml_content.push('\n');
}
}
// Write to file
fs::write(output_path, yaml_content)
.await
.context("Failed to write YAML file")?;
info!("Successfully exported YAML to: {}", output_path);
Ok(())
}
/// Generate consolidated report comparing all models
pub async fn generate_consolidated_report(&self, batch_id: Uuid) -> Result<String> {
let jobs = self.jobs.read().await;
let job = jobs
.get(&batch_id)
.ok_or_else(|| anyhow!("Batch job {} not found", batch_id))?;
let mut report = String::new();
// Header
report.push_str("╔════════════════════════════════════════════════════════════════╗\n");
report.push_str("║ BATCH TUNING CONSOLIDATED REPORT ║\n");
report.push_str("╚════════════════════════════════════════════════════════════════╝\n\n");
// Batch summary
report.push_str(&format!("Batch ID: {}\n", job.batch_id));
report.push_str(&format!("Status: {:?}\n", job.status));
report.push_str(&format!(
"Started: {}\n",
job.started_at.format("%Y-%m-%d %H:%M:%S UTC")
));
if let Some(completed_at) = job.completed_at {
let duration = completed_at.signed_duration_since(job.started_at);
report.push_str(&format!(
"Completed: {}\n",
completed_at.format("%Y-%m-%d %H:%M:%S UTC")
));
report.push_str(&format!("Duration: {} minutes\n", duration.num_minutes()));
}
report.push_str(&format!("\nModels Tuned: {}\n", job.results.len()));
report.push_str(&format!("Trials per Model: {}\n\n", job.trials_per_model));
// Per-model results
report.push_str("═══════════════════════════════════════════════════════════════\n");
report.push_str(" PER-MODEL RESULTS\n");
report.push_str("═══════════════════════════════════════════════════════════════\n\n");
for result in &job.results {
report.push_str(&format!("🔹 {}\n", result.model_type));
report.push_str(&format!(" Status: {:?}\n", result.status));
report.push_str(&format!(
" Trials Completed: {}\n",
result.trials_completed
));
if let Some(sharpe) = result.best_metrics.get("sharpe_ratio") {
report.push_str(&format!(" Best Sharpe Ratio: {:.4}\n", sharpe));
}
if let Some(loss) = result.best_metrics.get("training_loss") {
report.push_str(&format!(" Training Loss: {:.6}\n", loss));
}
if let Some(duration) = result.completed_at {
let model_duration = duration.signed_duration_since(result.started_at);
report.push_str(&format!(
" Duration: {} minutes\n",
model_duration.num_minutes()
));
}
if let Some(ref error) = result.error_message {
report.push_str(&format!(" ⚠️ Error: {}\n", error));
}
report.push('\n');
}
// Comparison table
let successful: Vec<_> = job
.results
.iter()
.filter(|r| r.status == TuningJobStatus::Completed)
.collect();
if !successful.is_empty() {
report.push_str("═══════════════════════════════════════════════════════════════\n");
report.push_str(" MODEL COMPARISON\n");
report.push_str("═══════════════════════════════════════════════════════════════\n\n");
report.push_str("┌──────────┬──────────────┬────────────────┐\n");
report.push_str("│ Model │ Sharpe Ratio │ Training Loss │\n");
report.push_str("├──────────┼──────────────┼────────────────┤\n");
for result in &successful {
let sharpe = result
.best_metrics
.get("sharpe_ratio")
.copied()
.unwrap_or(0.0);
let loss = result
.best_metrics
.get("training_loss")
.copied()
.unwrap_or(0.0);
report.push_str(&format!(
"{:<8}{:>12.4}{:>14.6}\n",
result.model_type, sharpe, loss
));
}
report.push_str("└──────────┴──────────────┴────────────────┘\n\n");
// Recommendation
if let Some(best) = successful.iter().max_by(|a, b| {
let a_sharpe = a.best_metrics.get("sharpe_ratio").copied().unwrap_or(0.0);
let b_sharpe = b.best_metrics.get("sharpe_ratio").copied().unwrap_or(0.0);
a_sharpe
.partial_cmp(&b_sharpe)
.unwrap_or(std::cmp::Ordering::Equal)
}) {
let best_sharpe = best
.best_metrics
.get("sharpe_ratio")
.copied()
.unwrap_or(0.0);
report.push_str("🏆 RECOMMENDATION\n");
report.push_str(&format!(
" Best Overall Model: {} (Sharpe Ratio: {:.4})\n",
best.model_type, best_sharpe
));
report.push_str(" Use these hyperparameters for production deployment.\n\n");
}
}
// YAML export info
report.push_str("═══════════════════════════════════════════════════════════════\n");
report.push_str(" EXPORT INFORMATION\n");
report.push_str("═══════════════════════════════════════════════════════════════\n\n");
report.push_str(&format!("YAML Export Path: {}\n", job.yaml_export_path));
if job.auto_export_yaml {
report.push_str("Status: ✅ Auto-exported\n");
} else {
report.push_str("Status: ⏳ Manual export required\n");
report.push_str(&format!(
" Command: tli tune batch export --batch-id {}\n",
job.batch_id
));
}
Ok(report)
}
/// Validate model type
fn is_valid_model(&self, model: &str) -> bool {
matches!(model, "DQN" | "PPO" | "MAMBA_2" | "TLOB" | "TFT" | "LIQUID")
}
}
#[cfg(test)]
#[allow(clippy::unwrap_used, clippy::expect_used)]
mod tests {
use super::*;
use crate::tuning_manager::TuningManager;
#[test]
fn test_dependency_resolution_simple() {
let working_dir = "/tmp/test_tuning".to_string();
let tuning_manager = Arc::new(TuningManager::new(
"/path/to/script.py".to_string(),
working_dir.clone(),
));
let manager = BatchTuningManager::new(tuning_manager, working_dir);
// TFT depends on MAMBA_2
let models = vec!["TFT".to_string(), "MAMBA_2".to_string()];
let resolved = manager.resolve_model_dependencies(&models);
assert_eq!(resolved.len(), 2);
let mamba_idx = resolved
.iter()
.position(|m| m == "MAMBA_2")
.expect("INVARIANT: MAMBA_2 must be in resolved models");
let tft_idx = resolved
.iter()
.position(|m| m == "TFT")
.expect("INVARIANT: TFT must be in resolved models");
assert!(mamba_idx < tft_idx, "MAMBA_2 must come before TFT");
}
#[test]
fn test_dependency_resolution_independent() {
let working_dir = "/tmp/test_tuning".to_string();
let tuning_manager = Arc::new(TuningManager::new(
"/path/to/script.py".to_string(),
working_dir.clone(),
)) as Arc<dyn TuningManagerTrait>;
let manager = BatchTuningManager::new(tuning_manager, working_dir);
// DQN and PPO are independent
let models = vec!["DQN".to_string(), "PPO".to_string()];
let resolved = manager.resolve_model_dependencies(&models);
assert_eq!(resolved.len(), 2);
assert!(resolved.contains(&"DQN".to_string()));
assert!(resolved.contains(&"PPO".to_string()));
}
#[test]
fn test_model_validation() {
let working_dir = "/tmp/test_tuning".to_string();
let tuning_manager = Arc::new(TuningManager::new(
"/path/to/script.py".to_string(),
working_dir.clone(),
)) as Arc<dyn TuningManagerTrait>;
let manager = BatchTuningManager::new(tuning_manager, working_dir);
assert!(manager.is_valid_model("DQN"));
assert!(manager.is_valid_model("PPO"));
assert!(manager.is_valid_model("MAMBA_2"));
assert!(manager.is_valid_model("TFT"));
assert!(!manager.is_valid_model("INVALID"));
}
}