fix: update evaluate_baseline for new backtest evaluator API

evaluate_dqn/evaluate_dqn_graphed now take (weights, branching_weights,
DqnBacktestConfig) instead of (weights, network_dims). This was the
compile error breaking all H100 CI runs.

Co-Authored-By: Claude Opus 4.6 (1M context) <noreply@anthropic.com>
This commit is contained in:
jgrusewski
2026-03-22 15:16:11 +01:00
parent f782e9f7ab
commit 39c636a0cb
2 changed files with 11 additions and 3 deletions

View File

@@ -1202,13 +1202,22 @@ fn evaluate_dqn_fold_gpu(
let weights = extract_dueling_weights(dqn.get_q_network_vars(), &stream)
.with_context(|| format!("extract_dueling_weights failed for fold {}", fold))?;
// Extract branching weights if branching DQN is active
use ml::cuda_pipeline::gpu_weights::extract_branching_weights;
let branching_weights = dqn.branching_q_network.as_ref()
.map(|br| extract_branching_weights(br.vars(), &stream))
.transpose()
.with_context(|| format!("extract_branching_weights failed for fold {}", fold))?;
let dqn_cfg = ml::cuda_pipeline::gpu_backtest_evaluator::DqnBacktestConfig::from_network_dims(network_dims);
if args.cuda_graphs {
info!(
" [DQN GPU] Using CUDA Graph-captured forward pass (fold {}, dims=({},{},{},{}))",
fold, network_dims.0, network_dims.1, network_dims.2, network_dims.3,
);
evaluator
.evaluate_dqn_graphed(&weights, network_dims)
.evaluate_dqn_graphed(&weights, branching_weights.as_ref(), &dqn_cfg)
.with_context(|| format!(
"GpuBacktestEvaluator::evaluate_dqn_graphed failed for fold {}", fold
))?
@@ -1218,7 +1227,7 @@ fn evaluate_dqn_fold_gpu(
fold, network_dims.0, network_dims.1, network_dims.2, network_dims.3,
);
evaluator
.evaluate_dqn(&weights, network_dims)
.evaluate_dqn(&weights, branching_weights.as_ref(), &dqn_cfg)
.with_context(|| format!(
"GpuBacktestEvaluator::evaluate_dqn failed for fold {}", fold
))?

View File

@@ -467,4 +467,3 @@ pub mod prelude {
}
// Tests for FactoredAction <-> TradingAction now live in ml-core::common::action
// trigger CI