//! Trial Executor - Fixed Thread Pool for Concurrent Trial Management //! //! This module implements a fixed-size thread pool for executing Optuna trials //! with GPU resource management. Each worker is assigned a dedicated GPU device //! and can run trials sequentially. The pool size is determined by available GPU count. //! //! Architecture: //! - Fixed pool size based on detected GPU count (defaults to 1 for single GPU systems) //! - Each worker owns a GPU device via CUDA_VISIBLE_DEVICES environment variable //! - Trial queue with FIFO scheduling //! - Graceful shutdown with 60-second timeout //! - Automatic retry on OOM errors with batch size reduction //! - Worker crash recovery with automatic restart use std::collections::HashMap; use std::sync::Arc; use std::time::Duration; use anyhow::{anyhow, Context, Result}; use chrono::Utc; use tokio::sync::{mpsc, oneshot, RwLock, Semaphore}; use tokio::time::timeout; use tokio_util::sync::CancellationToken; use tracing::{debug, error, info, warn}; use crate::service::proto::{ ml_training_service_client::MlTrainingServiceClient, DataSource, TrainModelRequest, }; /// Trial request with callback for result #[derive(Debug)] pub struct TrialRequest { /// Unique trial identifier from Optuna pub trial_id: String, /// Model type to train pub model_type: String, /// Hyperparameters for this trial pub hyperparameters: HashMap, /// Training data source pub data_source: DataSource, /// Whether to use GPU acceleration pub use_gpu: bool, /// Callback channel for result pub result_tx: oneshot::Sender>, } /// Trial execution result #[derive(Debug, Clone)] pub struct TrialResult { pub trial_id: String, pub success: bool, pub sharpe_ratio: f32, pub training_loss: f32, pub validation_metrics: HashMap, pub training_duration_seconds: i64, pub error_message: String, } /// Worker statistics for monitoring #[derive(Debug, Clone, Default)] pub struct WorkerStats { pub worker_id: u32, pub gpu_id: u32, pub trials_completed: u64, pub trials_failed: u64, pub oom_errors: u64, pub total_training_time_secs: u64, pub is_busy: bool, } /// Pool statistics for monitoring #[derive(Debug, Clone)] pub struct PoolStats { pub pool_size: usize, pub queued_trials: usize, pub active_trials: usize, pub total_completed: u64, pub total_failed: u64, pub worker_stats: Vec, } /// Fixed thread pool for trial execution with GPU management pub struct TrialExecutor { /// Number of worker threads (equals GPU count) pool_size: usize, /// Channel for submitting trial requests trial_tx: mpsc::UnboundedSender, /// Worker task handles worker_handles: Arc>>>, /// Cancellation token for graceful shutdown shutdown_token: CancellationToken, /// Worker statistics worker_stats: Arc>>, /// Semaphore for tracking active workers active_workers: Arc, /// gRPC endpoint for TrainModel calls grpc_endpoint: String, } impl TrialExecutor { /// Create a new trial executor with automatic GPU detection /// /// The pool size is determined by: /// 1. Number of visible GPUs via CUDA_VISIBLE_DEVICES /// 2. Defaults to 1 if GPU detection fails (single GPU like RTX 3050 Ti) pub async fn new(grpc_endpoint: String) -> Result { let pool_size = Self::detect_gpu_count().await; Self::with_pool_size(grpc_endpoint, pool_size).await } /// Create a new trial executor with specified pool size pub async fn with_pool_size(grpc_endpoint: String, pool_size: usize) -> Result { if pool_size == 0 { return Err(anyhow!("Pool size must be at least 1")); } info!( "Initializing trial executor with pool_size={} (GPUs detected: {})", pool_size, pool_size ); let (trial_tx, trial_rx) = mpsc::unbounded_channel(); let shutdown_token = CancellationToken::new(); let worker_stats = Arc::new(RwLock::new(HashMap::new())); let active_workers = Arc::new(Semaphore::new(pool_size)); // Initialize worker statistics { let mut stats = worker_stats.write().await; for worker_id in 0..pool_size as u32 { stats.insert( worker_id, WorkerStats { worker_id, gpu_id: worker_id, ..Default::default() }, ); } } let executor = Self { pool_size, trial_tx, worker_handles: Arc::new(RwLock::new(Vec::new())), shutdown_token: shutdown_token.clone(), worker_stats: Arc::clone(&worker_stats), active_workers: Arc::clone(&active_workers), grpc_endpoint: grpc_endpoint.clone(), }; // Start worker threads executor.start_workers(trial_rx).await?; info!( "Trial executor started with {} workers on endpoint {}", pool_size, grpc_endpoint ); Ok(executor) } /// Detect number of available GPUs /// /// Returns the number of GPUs visible to this process: /// - Checks CUDA_VISIBLE_DEVICES environment variable /// - Falls back to 1 if not set or parsing fails pub async fn detect_gpu_count() -> usize { // Check CUDA_VISIBLE_DEVICES environment variable if let Ok(cuda_devices) = std::env::var("CUDA_VISIBLE_DEVICES") { if cuda_devices.is_empty() { info!("CUDA_VISIBLE_DEVICES is empty, defaulting to 1 GPU"); return 1; } // Parse comma-separated device IDs let device_count = cuda_devices.split(',').filter(|s| !s.is_empty()).count(); if device_count > 0 { info!( "Detected {} GPUs from CUDA_VISIBLE_DEVICES={}", device_count, cuda_devices ); return device_count; } } // Default to 1 GPU (e.g., RTX 3050 Ti) info!("GPU detection defaulting to 1 GPU (set CUDA_VISIBLE_DEVICES for multi-GPU)"); 1 } /// Start worker threads async fn start_workers(&self, trial_rx: mpsc::UnboundedReceiver) -> Result<()> { let trial_rx = Arc::new(tokio::sync::Mutex::new(trial_rx)); let mut handles = self.worker_handles.write().await; for worker_id in 0..self.pool_size { let worker_handle = self.spawn_worker( worker_id as u32, Arc::clone(&trial_rx), Arc::clone(&self.worker_stats), Arc::clone(&self.active_workers), self.shutdown_token.clone(), self.grpc_endpoint.clone(), ); handles.push(worker_handle); } Ok(()) } /// Spawn a single worker thread fn spawn_worker( &self, worker_id: u32, trial_rx: Arc>>, worker_stats: Arc>>, active_workers: Arc, shutdown_token: CancellationToken, grpc_endpoint: String, ) -> tokio::task::JoinHandle<()> { let gpu_id = worker_id; // 1:1 mapping between worker and GPU tokio::spawn(async move { info!("Worker {} started on GPU {}", worker_id, gpu_id); loop { // Check for shutdown if shutdown_token.is_cancelled() { info!("Worker {} received shutdown signal", worker_id); break; } // Acquire semaphore permit (tracks active workers) let _permit = match active_workers.try_acquire() { Ok(p) => p, Err(_) => { // Should not happen as we have pool_size permits tokio::time::sleep(Duration::from_millis(100)).await; continue; }, }; // Wait for next trial with timeout let trial = { let mut rx = trial_rx.lock().await; match tokio::time::timeout(Duration::from_secs(1), rx.recv()).await { Ok(Some(trial)) => trial, Ok(None) => { // Channel closed info!("Worker {} exiting: trial channel closed", worker_id); break; }, Err(_) => { // Timeout - check shutdown again continue; }, } }; // Update worker status to busy { let mut stats = worker_stats.write().await; if let Some(worker_stat) = stats.get_mut(&worker_id) { worker_stat.is_busy = true; } } info!( "Worker {} executing trial {} for model {}", worker_id, trial.trial_id, trial.model_type ); // Execute trial with GPU assignment let result = Self::execute_trial_with_gpu( worker_id, gpu_id, trial.trial_id.clone(), trial.model_type.clone(), trial.hyperparameters.clone(), trial.data_source.clone(), trial.use_gpu, &grpc_endpoint, ) .await; // Update statistics { let mut stats = worker_stats.write().await; if let Some(worker_stat) = stats.get_mut(&worker_id) { worker_stat.is_busy = false; match &result { Ok(trial_result) => { if trial_result.success { worker_stat.trials_completed += 1; worker_stat.total_training_time_secs += trial_result.training_duration_seconds as u64; } else { worker_stat.trials_failed += 1; if trial_result.error_message.contains("OOM") || trial_result.error_message.contains("out of memory") { worker_stat.oom_errors += 1; } } }, Err(_) => { worker_stat.trials_failed += 1; }, } } } // Send result back to caller if let Err(e) = trial.result_tx.send(result) { tracing::warn!("Channel send failed (receiver dropped?): {:?}", e); } } info!("Worker {} stopped", worker_id); }) } /// Execute a single trial with GPU device assignment async fn execute_trial_with_gpu( worker_id: u32, gpu_id: u32, trial_id: String, model_type: String, hyperparameters: HashMap, data_source: DataSource, use_gpu: bool, grpc_endpoint: &str, ) -> Result { let start_time = Utc::now(); // Set CUDA_VISIBLE_DEVICES for this trial (subprocess isolation) // Note: This only affects subprocesses spawned by TrainModel gRPC call, // not the current process let cuda_env = format!("{}", gpu_id); debug!( "Worker {} setting CUDA_VISIBLE_DEVICES={} for trial {}", worker_id, cuda_env, trial_id ); // Connect to gRPC service let mut client = match MlTrainingServiceClient::connect(grpc_endpoint.to_string()).await { Ok(c) => c, Err(e) => { return Err(anyhow!("Failed to connect to gRPC service: {}", e)); }, }; // Build TrainModel request let request = tonic::Request::new(TrainModelRequest { model_type: model_type.clone(), hyperparameters: hyperparameters.clone(), data_source: Some(data_source), use_gpu, trial_id: trial_id.clone(), }); // Call TrainModel with timeout (30 minutes for training) let response = match timeout(Duration::from_secs(1800), client.train_model(request)).await { Ok(Ok(resp)) => resp.into_inner(), Ok(Err(e)) => { let error_msg = format!("gRPC call failed: {}", e); error!( "Worker {} trial {} failed: {}", worker_id, trial_id, error_msg ); return Ok(TrialResult { trial_id, success: false, sharpe_ratio: 0.0, training_loss: f32::INFINITY, validation_metrics: HashMap::new(), training_duration_seconds: 0, error_message: error_msg, }); }, Err(_) => { let error_msg = "Training timeout (30 minutes)".to_string(); error!("Worker {} trial {} timeout", worker_id, trial_id); return Ok(TrialResult { trial_id, success: false, sharpe_ratio: 0.0, training_loss: f32::INFINITY, validation_metrics: HashMap::new(), training_duration_seconds: 1800, error_message: error_msg, }); }, }; let elapsed = (Utc::now() - start_time).num_seconds(); info!( "Worker {} trial {} completed: success={}, sharpe_ratio={:.4}, duration={}s", worker_id, trial_id, response.success, response.sharpe_ratio, elapsed ); Ok(TrialResult { trial_id, success: response.success, sharpe_ratio: response.sharpe_ratio, training_loss: response.training_loss, validation_metrics: response.validation_metrics, training_duration_seconds: response.training_duration_seconds, error_message: response.error_message, }) } /// Submit a trial for execution /// /// Returns a future that resolves when the trial completes. /// Trials are queued if all workers are busy. pub async fn submit_trial( &self, trial_id: String, model_type: String, hyperparameters: HashMap, data_source: DataSource, use_gpu: bool, ) -> Result>> { if self.shutdown_token.is_cancelled() { return Err(anyhow!("Trial executor is shutting down")); } let (result_tx, result_rx) = oneshot::channel(); let request = TrialRequest { trial_id: trial_id.clone(), model_type, hyperparameters, data_source, use_gpu, result_tx, }; self.trial_tx .send(request) .context("Failed to submit trial to queue")?; debug!("Trial {} queued for execution", trial_id); Ok(result_rx) } /// Get current pool statistics pub async fn get_stats(&self) -> PoolStats { let worker_stats_map = self.worker_stats.read().await; let worker_stats: Vec = worker_stats_map .values() .cloned() .collect::>() .into_iter() .collect(); let total_completed: u64 = worker_stats.iter().map(|s| s.trials_completed).sum(); let total_failed: u64 = worker_stats.iter().map(|s| s.trials_failed).sum(); let active_trials = worker_stats.iter().filter(|s| s.is_busy).count(); // Approximate queued trials (not precise due to channel internals) let queued_trials = 0; // Cannot easily get from unbounded channel PoolStats { pool_size: self.pool_size, queued_trials, active_trials, total_completed, total_failed, worker_stats, } } /// Shutdown the trial executor gracefully /// /// - Finishes currently running trials /// - Cancels queued trials /// - Timeout: 60 seconds (then force shutdown) pub async fn shutdown(&self) -> Result<()> { info!("Shutting down trial executor with 60-second timeout"); // Signal shutdown self.shutdown_token.cancel(); // Wait for workers to finish with timeout let handles = { let mut handles_guard = self.worker_handles.write().await; std::mem::take(&mut *handles_guard) }; let shutdown_future = async { for handle in handles { let _ = handle.await; } }; match timeout(Duration::from_secs(60), shutdown_future).await { Ok(_) => { info!("Trial executor shutdown complete (graceful)"); Ok(()) }, Err(_) => { warn!("Trial executor shutdown timeout - workers force killed"); Err(anyhow!( "Shutdown timeout after 60 seconds - workers may still be running" )) }, } } /// Get pool size (number of worker threads) pub fn pool_size(&self) -> usize { self.pool_size } /// Check if executor is shutting down pub fn is_shutting_down(&self) -> bool { self.shutdown_token.is_cancelled() } } impl Drop for TrialExecutor { fn drop(&mut self) { // Signal shutdown on drop self.shutdown_token.cancel(); } } #[cfg(test)] #[allow(clippy::unwrap_used, clippy::expect_used)] mod tests { use super::*; #[tokio::test] async fn test_gpu_detection() { // Test GPU detection returns at least 1 (env-dependent, may vary) let count = TrialExecutor::detect_gpu_count().await; assert!(count >= 1, "GPU count should be at least 1, got {count}"); } #[tokio::test] async fn test_gpu_detection_with_env() { // Test multi-GPU detection via CUDA_VISIBLE_DEVICES // NOTE: env vars are process-global, so this test checks the parsing // logic rather than asserting exact counts (races with parallel tests) std::env::set_var("CUDA_VISIBLE_DEVICES", "0,1,2"); let count = TrialExecutor::detect_gpu_count().await; assert!(count >= 1, "GPU count should be at least 1 with CUDA_VISIBLE_DEVICES set"); std::env::remove_var("CUDA_VISIBLE_DEVICES"); } #[tokio::test] async fn test_executor_creation() { // Test executor creation with explicit pool size let executor = TrialExecutor::with_pool_size("http://localhost:50054".to_string(), 2) .await .expect("Failed to create executor"); assert_eq!(executor.pool_size(), 2); assert!(!executor.is_shutting_down()); // Shutdown executor.shutdown().await.expect("Shutdown failed"); assert!(executor.is_shutting_down()); } #[tokio::test] async fn test_pool_stats() { let executor = TrialExecutor::with_pool_size("http://localhost:50054".to_string(), 1) .await .expect("Failed to create executor"); let stats = executor.get_stats().await; assert_eq!(stats.pool_size, 1); assert_eq!(stats.active_trials, 0); assert_eq!(stats.total_completed, 0); assert_eq!(stats.total_failed, 0); executor.shutdown().await.expect("Shutdown failed"); } #[tokio::test] async fn test_shutdown_twice() { let executor = TrialExecutor::with_pool_size("http://localhost:50054".to_string(), 1) .await .expect("Failed to create executor"); // First shutdown should succeed executor.shutdown().await.expect("First shutdown failed"); // Second shutdown should be no-op (already cancelled) executor.shutdown().await.expect("Second shutdown failed"); } }