diff --git a/crates/ml/examples/hyperopt_baseline_rl.rs b/crates/ml/examples/hyperopt_baseline_rl.rs index 1d4fb4fdb..06ee835a4 100644 --- a/crates/ml/examples/hyperopt_baseline_rl.rs +++ b/crates/ml/examples/hyperopt_baseline_rl.rs @@ -124,12 +124,14 @@ struct Args { fn build_model_result( best_objective: f64, best_params_json: Value, + top_k_params: Vec, num_trials: usize, elapsed_secs: f64, ) -> Value { serde_json::json!({ "best_objective": best_objective, "best_params": best_params_json, + "top_k_params": top_k_params, "trials": num_trials, "elapsed_secs": elapsed_secs, }) @@ -209,9 +211,31 @@ fn run_dqn_hyperopt(args: &Args, parallel: usize, gpu_devices: &[candle_core::De } }; + // Extract top-5 trials sorted by objective (ascending = best first) + let mut sorted_trials = result.all_trials.clone(); + sorted_trials.sort_by(|a, b| { + a.objective + .partial_cmp(&b.objective) + .unwrap_or(std::cmp::Ordering::Equal) + }); + let top_k: Vec = sorted_trials + .iter() + .take(5) + .filter_map(|t| { + serde_json::to_value(&t.params).ok().map(|params| { + serde_json::json!({ + "params": params, + "objective": t.objective, + "trial_num": t.trial_num, + }) + }) + }) + .collect(); + Ok(build_model_result( result.best_objective, best_params_json, + top_k, result.all_trials.len(), elapsed, )) @@ -290,9 +314,31 @@ fn run_ppo_hyperopt(args: &Args, parallel: usize, gpu_devices: &[candle_core::De } }; + // Extract top-5 trials sorted by objective (ascending = best first) + let mut sorted_trials = result.all_trials.clone(); + sorted_trials.sort_by(|a, b| { + a.objective + .partial_cmp(&b.objective) + .unwrap_or(std::cmp::Ordering::Equal) + }); + let top_k: Vec = sorted_trials + .iter() + .take(5) + .filter_map(|t| { + serde_json::to_value(&t.params).ok().map(|params| { + serde_json::json!({ + "params": params, + "objective": t.objective, + "trial_num": t.trial_num, + }) + }) + }) + .collect(); + Ok(build_model_result( result.best_objective, best_params_json, + top_k, result.all_trials.len(), elapsed, )) diff --git a/crates/ml/examples/train_baseline_rl.rs b/crates/ml/examples/train_baseline_rl.rs index 87315885c..a9f26c569 100644 --- a/crates/ml/examples/train_baseline_rl.rs +++ b/crates/ml/examples/train_baseline_rl.rs @@ -134,6 +134,12 @@ struct Args { /// Half-spread slippage is added to `tx_cost_bps` per trade. #[arg(long, default_value_t = 1.0)] spread_ticks: f64, + + /// Number of top hyperopt configs to train as ensemble (default 1 = best only). + /// When > 1, loads `top_k_params` from the hyperopt JSON and trains a separate + /// model for each param set, saving as `dqn_ensemble_{k}_fold_{fold}.safetensors`. + #[arg(long, default_value_t = 1)] + ensemble_top_k: usize, } // --------------------------------------------------------------------------- @@ -176,6 +182,41 @@ fn load_hyperopt_params(hp_path: &Option, model_key: &str) -> Option, model_key: &str, k: usize) -> Vec> { + let file_path = match hp_path.as_ref() { + Some(p) if p.exists() => p, + _ => return vec![None; k], + }; + let Ok(contents) = std::fs::read_to_string(file_path) else { + return vec![None; k]; + }; + let Ok(json): Result = serde_json::from_str(&contents) else { + return vec![None; k]; + }; + let top_k = json + .get(model_key) + .and_then(|m| m.get("top_k_params")) + .and_then(|v| v.as_array()); + if let Some(arr) = top_k { + arr.iter() + .take(k) + .map(|entry| entry.get("params").cloned()) + .collect() + } else { + // Fallback: just use best_params as the single entry + let best = json + .get(model_key) + .and_then(|m| m.get("best_params")) + .cloned(); + vec![best] + } +} + fn hp_f64(params: &Option, key: &str) -> Option { params.as_ref()?.get(key)?.as_f64() } @@ -296,6 +337,7 @@ fn train_dqn_fold( output_dir: &Path, hp: &Option, pre_uploaded_gpu_data: Option, + checkpoint_prefix: &str, ) -> Result { // Caller must pre-align bars to features (skip warmup period before calling). debug_assert_eq!( @@ -416,9 +458,10 @@ fn train_dqn_fold( // Checkpoint callback: save best model to output directory let output_dir_owned = output_dir.to_path_buf(); + let prefix_owned = checkpoint_prefix.to_owned(); let checkpoint_callback = move |epoch: usize, data: Vec, is_best: bool| -> Result { let suffix = if is_best { "best" } else { &format!("epoch{}", epoch) }; - let ckpt_path = output_dir_owned.join(format!("dqn_fold{}_{}.safetensors", fold, suffix)); + let ckpt_path = output_dir_owned.join(format!("{}_fold{}_{}.safetensors", prefix_owned, fold, suffix)); // Atomic write: write to .tmp then rename (POSIX rename is atomic) let tmp_path = ckpt_path.with_extension("safetensors.tmp"); std::fs::write(&tmp_path, &data) @@ -427,7 +470,7 @@ fn train_dqn_fold( drop(std::fs::remove_file(&tmp_path)); return Err(e).with_context(|| format!("Failed to rename checkpoint: {} -> {}", tmp_path.display(), ckpt_path.display())); } - info!(" [DQN] Fold {} saved checkpoint: {}", fold, ckpt_path.display()); + info!(" [DQN] Fold {} saved checkpoint: {} (prefix: {})", fold, ckpt_path.display(), prefix_owned); Ok(ckpt_path.to_string_lossy().into_owned()) }; @@ -617,6 +660,9 @@ fn run_training(args: &Args) -> Result> { if let Some(ref hp_path) = args.hyperopt_params { info!(" Hyperopt params: {}", hp_path.display()); } + if args.ensemble_top_k > 1 { + info!(" Ensemble top-K: {} (training multiple models per fold)", args.ensemble_top_k); + } // 1. Load all OHLCV bars from DBN files info!("Step 1/5: Loading OHLCV bars from DBN files..."); @@ -737,34 +783,90 @@ fn run_training(args: &Args) -> Result> { // Train DQN if train_dqn { - let hp = load_hyperopt_params(&args.hyperopt_params, "dqn"); let fold_str = fold_idx.to_string(); let fold_start = std::time::Instant::now(); // Take pre-uploaded GPU data from previous fold's background upload - let gpu_data_for_fold = dqn_gpu_staged.take(); + let mut gpu_data_for_fold = dqn_gpu_staged.take(); - match train_dqn_fold( - window.fold, - &train_norm, - &val_norm, - &train_bars_aligned, - &val_bars_aligned, - args, - &args.output_dir, - &hp, - gpu_data_for_fold, - ) { - Ok(best_loss) => { - let elapsed = fold_start.elapsed().as_secs_f64(); - metrics::set_epoch("dqn", &fold_str, fold_idx as f64); - metrics::set_epoch_loss("dqn", &fold_str, best_loss); - metrics::set_validation_loss("dqn", &fold_str, best_loss); - metrics::set_iteration_seconds("dqn", &fold_str, elapsed); - dqn_results.push((window.fold, best_loss)); + if args.ensemble_top_k > 1 && args.hyperopt_params.is_some() { + // Ensemble mode: train one model per top-K hyperopt param set + let param_sets = load_top_k_params( + &args.hyperopt_params, + "dqn", + args.ensemble_top_k, + ); + for (k, hp) in param_sets.iter().enumerate() { + info!( + " [DQN] Training ensemble member {}/{} on fold {}", + k + 1, + param_sets.len(), + window.fold + ); + let prefix = format!("dqn_ensemble_{}", k); + // Only the first ensemble member uses the pre-uploaded GPU data + let gpu_data = if k == 0 { gpu_data_for_fold.take() } else { None }; + match train_dqn_fold( + window.fold, + &train_norm, + &val_norm, + &train_bars_aligned, + &val_bars_aligned, + args, + &args.output_dir, + hp, + gpu_data, + &prefix, + ) { + Ok(best_loss) => { + info!( + " [DQN] Ensemble member {} fold {} best_loss={:.6}", + k, window.fold, best_loss + ); + // Record the best ensemble member's loss as the fold result + if k == 0 { + let elapsed = fold_start.elapsed().as_secs_f64(); + metrics::set_epoch("dqn", &fold_str, fold_idx as f64); + metrics::set_epoch_loss("dqn", &fold_str, best_loss); + metrics::set_validation_loss("dqn", &fold_str, best_loss); + metrics::set_iteration_seconds("dqn", &fold_str, elapsed); + dqn_results.push((window.fold, best_loss)); + } + } + Err(e) => { + error!( + " [DQN] Ensemble member {} fold {} failed: {}", + k, window.fold, e + ); + } + } } - Err(e) => { - error!(" [DQN] Fold {} failed: {}", window.fold, e); + } else { + // Single-model mode (default) + let hp = load_hyperopt_params(&args.hyperopt_params, "dqn"); + match train_dqn_fold( + window.fold, + &train_norm, + &val_norm, + &train_bars_aligned, + &val_bars_aligned, + args, + &args.output_dir, + &hp, + gpu_data_for_fold, + "dqn", + ) { + Ok(best_loss) => { + let elapsed = fold_start.elapsed().as_secs_f64(); + metrics::set_epoch("dqn", &fold_str, fold_idx as f64); + metrics::set_epoch_loss("dqn", &fold_str, best_loss); + metrics::set_validation_loss("dqn", &fold_str, best_loss); + metrics::set_iteration_seconds("dqn", &fold_str, elapsed); + dqn_results.push((window.fold, best_loss)); + } + Err(e) => { + error!(" [DQN] Fold {} failed: {}", window.fold, e); + } } } } diff --git a/crates/ml/src/hyperopt/adapters/dqn.rs b/crates/ml/src/hyperopt/adapters/dqn.rs index 7d3f906dc..200ce1dad 100644 --- a/crates/ml/src/hyperopt/adapters/dqn.rs +++ b/crates/ml/src/hyperopt/adapters/dqn.rs @@ -2850,173 +2850,232 @@ impl HyperparameterOptimizable for DQNTrainer { tracing::warn!("No validation data available for backtest"); None } else { - // ── Chunked batch inference backtest ────────────────────────── + // ── Multi-window chunked batch inference backtest ─────────── // - // Process bars in chunks of EVAL_CHUNK_SIZE (1024). Each chunk: - // 1. Extracts state vectors on CPU (with current portfolio features) - // 2. Single GPU forward pass → all Q-values at once - // 3. Sequential trade simulation on CPU - // 4. Syncs portfolio state to trainer for next chunk + // Split validation data into `window_count` non-overlapping + // windows and run an independent backtest on each. The final + // objective uses `mean(Sharpe) - 0.5 * std(Sharpe)` to penalize + // inconsistency across windows, reducing overfit to one segment. // - // vs. the old per-bar loop: ~1000× fewer GPU kernel launches. + // Within each window the same chunked GPU inference strategy is + // used: EVAL_CHUNK_SIZE bars per GPU forward pass → sequential + // trade simulation on CPU → portfolio sync between chunks. const EVAL_CHUNK_SIZE: usize = 1024; + const WINDOW_COUNT: usize = 3; let kelly_fraction = internal_trainer.get_kelly_fraction(); #[allow(clippy::cast_possible_truncation)] let eval_capital = self.initial_capital as f32; - let mut engine = EvaluationEngine::new_with_kelly(eval_capital, kelly_fraction); let agent_arc = internal_trainer.get_agent().clone(); let device = internal_trainer.device().clone(); let total_bars = val_close_prices.len(); - let num_chunks = total_bars.div_ceil(EVAL_CHUNK_SIZE); - let mut ohlcv_bars = Vec::with_capacity(total_bars); - let mut conversion_failures: usize = 0; - let mut unique_actions = std::collections::HashSet::new(); + let window_size = total_bars / WINDOW_COUNT; - tracing::info!( - "Batched backtest: {} bars in {} chunks of {} (device: {:?})", - total_bars, num_chunks, EVAL_CHUNK_SIZE, device, - ); + // Need at least 1 bar per window to produce meaningful metrics + if window_size == 0 { + tracing::warn!( + "Validation data too small for multi-window backtest ({} bars < {} windows)", + total_bars, WINDOW_COUNT, + ); + None + } else { + tracing::info!( + "Multi-window backtest: {} windows x {} bars (total={}, device={:?})", + WINDOW_COUNT, window_size, total_bars, device, + ); - // Acquire read lock once — batch_softmax_actions is &self (no mutation) - let runtime = tokio::runtime::Runtime::new().map_err(|e| { - MLError::TrainingError(format!("Failed to create runtime for backtest: {}", e)) - })?; - let agent_guard = runtime.block_on(agent_arc.read()); + // Acquire read lock once — batch_hierarchical_softmax_actions is &self + let runtime = tokio::runtime::Runtime::new().map_err(|e| { + MLError::TrainingError(format!("Failed to create runtime for backtest: {}", e)) + })?; + let agent_guard = runtime.block_on(agent_arc.read()); - for chunk_idx in 0..num_chunks { - let chunk_start = chunk_idx * EVAL_CHUNK_SIZE; - let chunk_end = (chunk_start + EVAL_CHUNK_SIZE).min(total_bars); - let chunk_len = chunk_end - chunk_start; + // Collect per-window metrics + let mut window_metrics: Vec = Vec::with_capacity(WINDOW_COUNT); + let mut all_unique_actions: std::collections::HashSet = std::collections::HashSet::new(); + let mut total_conversion_failures: usize = 0; - // 1. Extract state vectors on CPU (portfolio features from trainer) - // Re-borrow val_data per chunk to allow mutable portfolio updates between chunks - let mut flat_states: Vec = Vec::with_capacity(chunk_len * 54); - let mut valid_mask: Vec = Vec::with_capacity(chunk_len); + #[allow(clippy::indexing_slicing)] // window indices verified by window_size > 0 + for win_idx in 0..WINDOW_COUNT { + let win_start = win_idx * window_size; + // Last window absorbs any remainder bars + let win_end = if win_idx + 1 == WINDOW_COUNT { + total_bars + } else { + (win_idx + 1) * window_size + }; + let win_len = win_end - win_start; + let win_prices = &val_close_prices[win_start..win_end]; - { - let val_data = internal_trainer.get_val_data(); - let chunk = &val_data[chunk_start..chunk_end]; + // Reset portfolio state for this window + internal_trainer.set_portfolio_for_backtest(0.0, 0.0, 0.0); + let mut engine = EvaluationEngine::new_with_kelly(eval_capital, kelly_fraction); - for (i, (feature_vec, _target)) in chunk.iter().enumerate() { - let close_price = val_close_prices[chunk_start + i]; - match internal_trainer.convert_to_state_vec(feature_vec, close_price) { - Ok(sv) => { - flat_states.extend_from_slice(&sv); - valid_mask.push(true); + let num_chunks = win_len.div_ceil(EVAL_CHUNK_SIZE); + let mut ohlcv_bars = Vec::with_capacity(win_len); + let mut conversion_failures: usize = 0; + let mut unique_actions = std::collections::HashSet::new(); + + for chunk_idx in 0..num_chunks { + let chunk_start = chunk_idx * EVAL_CHUNK_SIZE; + let chunk_end = (chunk_start + EVAL_CHUNK_SIZE).min(win_len); + let chunk_len = chunk_end - chunk_start; + + // Absolute indices into the full val_data / val_close_prices + let abs_start = win_start + chunk_start; + let abs_end = win_start + chunk_end; + + // 1. Extract state vectors on CPU + let mut flat_states: Vec = Vec::with_capacity(chunk_len * 54); + let mut valid_mask: Vec = Vec::with_capacity(chunk_len); + + { + let val_data = internal_trainer.get_val_data(); + let chunk = &val_data[abs_start..abs_end]; + + for (i, (feature_vec, _target)) in chunk.iter().enumerate() { + let close_price = val_close_prices[abs_start + i]; + match internal_trainer.convert_to_state_vec(feature_vec, close_price) { + Ok(sv) => { + flat_states.extend_from_slice(&sv); + valid_mask.push(true); + } + Err(_) => { + conversion_failures += 1; + flat_states.extend(std::iter::repeat_n(0.0_f32, 54)); + valid_mask.push(false); + } + } } - Err(_) => { - conversion_failures += 1; - flat_states.extend(std::iter::repeat_n(0.0_f32, 54)); - valid_mask.push(false); + } // val_data borrow dropped + + // 2. Single GPU forward pass for entire chunk + let batch_tensor = + candle_core::Tensor::from_slice(&flat_states, (chunk_len, 54), &device) + .map_err(|e| { + MLError::ModelError(format!( + "Batch tensor creation failed: {}", + e + )) + })?; + + let action_indices = + agent_guard.batch_hierarchical_softmax_actions(&batch_tensor, params.eval_softmax_temp)?; + + // 3. Sequential trade simulation on CPU + for (i, &action_idx) in action_indices.iter().enumerate() { + if !valid_mask[i] { + continue; } + let bar_idx = chunk_start + i; // window-local index + let close = win_prices[bar_idx] as f32; + + unique_actions.insert(action_idx); + let factored = crate::dqn::FactoredAction::from_index(action_idx)?; + + let bar = OHLCVBarF32 { + timestamp: (abs_start + i) as i64, + open: close, + high: close, + low: close, + close, + volume: 0.0, + }; + engine.process_bar_factored(bar_idx, &bar, &factored); + ohlcv_bars.push(bar); + } + + // 4. Sync portfolio state to trainer for next chunk's features + if chunk_idx + 1 < num_chunks { + let pos_size = engine.current_exposure as f32; + let entry_price = engine.exposure_entry_price; + let last_close = win_prices.get(chunk_end - 1) + .copied() + .unwrap_or(0.0) as f32; + internal_trainer.set_portfolio_for_backtest( + pos_size, + entry_price, + last_close, + ); } } - } // val_data borrow dropped here - // 2. Single GPU forward pass for entire chunk - let batch_tensor = - candle_core::Tensor::from_slice(&flat_states, (chunk_len, 54), &device) - .map_err(|e| { - MLError::ModelError(format!( - "Batch tensor creation failed: {}", - e - )) - })?; - - // Hierarchical factored softmax: first picks exposure level - // (Short100..Long100), then order×urgency within it. Prevents - // collapse to a single exposure bucket that produces no trades. - let action_indices = - agent_guard.batch_hierarchical_softmax_actions(&batch_tensor, params.eval_softmax_temp)?; - - // 3. Sequential trade simulation on CPU - for (i, &action_idx) in action_indices.iter().enumerate() { - if !valid_mask[i] { - continue; + // Close any open factored exposure at window end + if let Some(last_bar) = ohlcv_bars.last() { + engine.close_factored_position(ohlcv_bars.len() - 1, last_bar); } - let bar_idx = chunk_start + i; - let close = val_close_prices[bar_idx] as f32; - unique_actions.insert(action_idx); - let factored = crate::dqn::FactoredAction::from_index(action_idx)?; - - let bar = OHLCVBarF32 { - timestamp: bar_idx as i64, - open: close, - high: close, - low: close, - close, - volume: 0.0, - }; - engine.process_bar_factored(bar_idx, &bar, &factored); - ohlcv_bars.push(bar); - } - - // 4. Sync portfolio state to trainer for next chunk's features - if chunk_idx + 1 < num_chunks { - let pos_size = engine.current_exposure as f32; - let entry_price = engine.exposure_entry_price; - let last_close = val_close_prices.get(chunk_end - 1) - .copied() - .unwrap_or(0.0) as f32; - internal_trainer.set_portfolio_for_backtest( - pos_size, - entry_price, - last_close, + let wm = PerformanceMetrics::from_trades( + &engine.trades, + engine.initial_capital, + &ohlcv_bars, ); + + tracing::info!( + "Window {}/{}: {} bars, {} trades, Sharpe {:.4}, WinRate {:.2}%, \ + MaxDD {:.2}%, Return {:.2}%, unique_actions={}/{}, conv_fail={}", + win_idx + 1, WINDOW_COUNT, ohlcv_bars.len(), wm.total_trades, + wm.sharpe_ratio, wm.win_rate, wm.max_drawdown_pct, + wm.total_return_pct, unique_actions.len(), 45, conversion_failures, + ); + + all_unique_actions.extend(&unique_actions); + total_conversion_failures += conversion_failures; + window_metrics.push(wm); } + + // Release read lock + drop(agent_guard); + + // ── Aggregate across windows ──────────────────────────── + let n = window_metrics.len() as f64; + let sharpes: Vec = window_metrics.iter().map(|m| m.sharpe_ratio).collect(); + let mean_sharpe = sharpes.iter().sum::() / n; + let std_sharpe = (sharpes.iter().map(|s| (s - mean_sharpe).powi(2)).sum::() / n).sqrt(); + // Penalize inconsistency: configs that work across all windows get rewarded + let adjusted_sharpe = mean_sharpe - 0.5 * std_sharpe; + + let mean_win_rate = window_metrics.iter().map(|m| m.win_rate).sum::() / n; + let max_drawdown = window_metrics.iter().map(|m| m.max_drawdown_pct) + .fold(0.0_f64, f64::max); + let mean_return = window_metrics.iter().map(|m| m.total_return_pct).sum::() / n; + let total_trades: usize = window_metrics.iter().map(|m| m.total_trades).sum(); + let mean_sortino = window_metrics.iter().map(|m| m.sortino_ratio).sum::() / n; + let mean_calmar = window_metrics.iter().map(|m| m.calmar_ratio).sum::() / n; + let mean_var_95 = window_metrics.iter().map(|m| m.var_95).sum::() / n; + let mean_cvar_95 = window_metrics.iter().map(|m| m.cvar_95).sum::() / n; + let mean_beta = window_metrics.iter().map(|m| m.beta).sum::() / n; + let mean_alpha = window_metrics.iter().map(|m| m.alpha).sum::() / n; + let mean_info_ratio = window_metrics.iter().map(|m| m.information_ratio).sum::() / n; + let mean_omega = window_metrics.iter().map(|m| m.omega_ratio).sum::() / n; + + tracing::info!( + "Multi-window backtest: {} windows x {} bars, adjusted_sharpe={:.4} \ + (mean={:.4}, std={:.4}), total_trades={}, unique_actions={}/{}, \ + conversion_failures={}", + WINDOW_COUNT, window_size, adjusted_sharpe, mean_sharpe, std_sharpe, + total_trades, all_unique_actions.len(), 45, total_conversion_failures, + ); + + Some(BacktestMetrics { + sharpe_ratio: adjusted_sharpe, + win_rate: mean_win_rate, + max_drawdown_pct: max_drawdown, + total_return_pct: mean_return, + total_trades, + sortino_ratio: mean_sortino, + calmar_ratio: mean_calmar, + var_95: mean_var_95, + cvar_95: mean_cvar_95, + beta: mean_beta, + alpha: mean_alpha, + information_ratio: mean_info_ratio, + omega_ratio: mean_omega, + unique_actions: all_unique_actions.len(), + }) } - - // Release read lock - drop(agent_guard); - - // Close any open factored exposure at end - if let Some(last_bar) = ohlcv_bars.last() { - engine.close_factored_position(ohlcv_bars.len() - 1, last_bar); - } - - // Calculate performance metrics - let metrics = PerformanceMetrics::from_trades( - &engine.trades, - engine.initial_capital, - &ohlcv_bars, - ); - - tracing::info!( - "Batched backtest complete: {} bars, {} chunks, {} trades, \ - Sharpe {:.4}, Win Rate {:.2}%, Max DD {:.2}%, Return {:.2}%, \ - unique_actions={}/{}, conversion_failures={}", - ohlcv_bars.len(), - num_chunks, - metrics.total_trades, - metrics.sharpe_ratio, - metrics.win_rate, - metrics.max_drawdown_pct, - metrics.total_return_pct, - unique_actions.len(), - 45, // 5 exposure × 3 order × 3 urgency - conversion_failures, - ); - - Some(BacktestMetrics { - sharpe_ratio: metrics.sharpe_ratio, - win_rate: metrics.win_rate, - max_drawdown_pct: metrics.max_drawdown_pct, - total_return_pct: metrics.total_return_pct, - total_trades: metrics.total_trades, - sortino_ratio: metrics.sortino_ratio, - calmar_ratio: metrics.calmar_ratio, - var_95: metrics.var_95, - cvar_95: metrics.cvar_95, - beta: metrics.beta, - alpha: metrics.alpha, - information_ratio: metrics.information_ratio, - omega_ratio: metrics.omega_ratio, - unique_actions: unique_actions.len(), - }) } } else { None diff --git a/infra/k8s/argo/training-workflow-template.yaml b/infra/k8s/argo/training-workflow-template.yaml index c9788daba..c86de2b48 100644 --- a/infra/k8s/argo/training-workflow-template.yaml +++ b/infra/k8s/argo/training-workflow-template.yaml @@ -45,6 +45,8 @@ spec: value: "1.0" - name: initial-capital value: "35000" + - name: ensemble-top-k + value: "5" - name: cuda-compute-cap value: "90" @@ -304,7 +306,8 @@ spec: --data-dir {{workflow.parameters.data-dir}} \ --output-dir /workspace/output \ --max-steps-per-epoch {{workflow.parameters.train-epochs}} \ - $HYPEROPT_FLAG + $HYPEROPT_FLAG \ + --ensemble-top-k {{workflow.parameters.ensemble-top-k}} echo "=== Training complete ===" ls -lh /workspace/output/ @@ -377,6 +380,7 @@ spec: HYPEROPT_FLAG="--hyperopt-params $HYPEROPT_FILE" fi + # Evaluate all ensemble members (evaluator auto-discovers *_ensemble_*.safetensors) echo "Evaluating $MODEL" evaluate_baseline \ --model "$MODEL" \