From 3a6cfe9a9bf79e5f8e10dda116a050c019e4bc3a Mon Sep 17 00:00:00 2001 From: jgrusewski Date: Wed, 15 Apr 2026 00:36:34 +0200 Subject: [PATCH] fix: correct CUDA kernel names for VSN/GLU in constructor MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit vsn_bottleneck_fwd → variable_select_bottleneck glu_gate_combine → glu_combine glu_gate_backward → glu_backward 19/19 smoke tests pass with full CPBI stack. Co-Authored-By: Claude Opus 4.6 (1M context) --- crates/ml/src/cuda_pipeline/gpu_dqn_trainer.rs | 12 ++++++------ 1 file changed, 6 insertions(+), 6 deletions(-) diff --git a/crates/ml/src/cuda_pipeline/gpu_dqn_trainer.rs b/crates/ml/src/cuda_pipeline/gpu_dqn_trainer.rs index eb6a876ea..8dc152265 100644 --- a/crates/ml/src/cuda_pipeline/gpu_dqn_trainer.rs +++ b/crates/ml/src/cuda_pipeline/gpu_dqn_trainer.rs @@ -2769,12 +2769,12 @@ impl GpuDqnTrainer { .map_err(|e| MLError::ModelError(format!("selectivity_forward load: {e}")))?; let sel_bwd_kernel = cpbi_module.load_function("selectivity_backward") .map_err(|e| MLError::ModelError(format!("selectivity_backward load: {e}")))?; - let vsn_kernel = cpbi_module.load_function("vsn_bottleneck_fwd") - .map_err(|e| MLError::ModelError(format!("vsn_bottleneck_fwd load: {e}")))?; - let glu_combine_kernel = cpbi_module.load_function("glu_gate_combine") - .map_err(|e| MLError::ModelError(format!("glu_gate_combine load: {e}")))?; - let glu_backward_kernel = cpbi_module.load_function("glu_gate_backward") - .map_err(|e| MLError::ModelError(format!("glu_gate_backward load: {e}")))?; + let vsn_kernel = cpbi_module.load_function("variable_select_bottleneck") + .map_err(|e| MLError::ModelError(format!("variable_select_bottleneck load: {e}")))?; + let glu_combine_kernel = cpbi_module.load_function("glu_combine") + .map_err(|e| MLError::ModelError(format!("glu_combine load: {e}")))?; + let glu_backward_kernel = cpbi_module.load_function("glu_backward") + .map_err(|e| MLError::ModelError(format!("glu_backward load: {e}")))?; info!("GpuDqnTrainer: Q-attn + selectivity + VSN + GLU kernels loaded"); // ── Compile CQL penalty kernel (if enabled) ──────────────────────