From bd8b84a2a76e958f718eb4454220c071dccc3298 Mon Sep 17 00:00:00 2001 From: jgrusewski Date: Wed, 15 Apr 2026 23:52:24 +0200 Subject: [PATCH] =?UTF-8?q?fix:=20ensemble=5Faggregate=5Fkernel=20OOB=20?= =?UTF-8?q?=E2=80=94=20buffers=20sized=20for=20total=5Factions(12)=20not?= =?UTF-8?q?=20num=5Fatoms(51)?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit ensemble_mean_q_buf and ensemble_var_q_buf were allocated as batch_size * total_actions (12), but the kernel writes batch_size * num_atoms (51) elements. 2977 OOB write errors. Fixed: allocate batch_size * num_atoms. compute-sanitizer: 0 errors. Co-Authored-By: Claude Opus 4.6 (1M context) --- crates/ml/src/trainers/dqn/fused_training.rs | 12 +++++++----- 1 file changed, 7 insertions(+), 5 deletions(-) diff --git a/crates/ml/src/trainers/dqn/fused_training.rs b/crates/ml/src/trainers/dqn/fused_training.rs index b60d41530..d02917f0a 100644 --- a/crates/ml/src/trainers/dqn/fused_training.rs +++ b/crates/ml/src/trainers/dqn/fused_training.rs @@ -244,10 +244,10 @@ pub(crate) struct FusedTrainingCtx { /// F32 because output layer GemmEx writes f32 (no f32 truncation overflow). /// None when ensemble_count <= 1. pub(crate) ensemble_logits_buf: Option>, - /// Pre-allocated buffer: [B * total_actions] for mean Q-values across K heads. + /// Pre-allocated buffer: [B * num_atoms] for mean Q-values across K heads. /// None when ensemble_count <= 1. pub(crate) ensemble_mean_q_buf: Option>, - /// Pre-allocated buffer: [B * total_actions] for Q-value variance across K heads. + /// Pre-allocated buffer: [B * num_atoms] for Q-value variance across K heads. /// None when ensemble_count <= 1. pub(crate) ensemble_var_q_buf: Option>, /// Pre-allocated buffer: [1] for accumulated diversity loss scalar. @@ -596,12 +596,14 @@ impl FusedTrainingCtx { // Pre-allocate ensemble buffers. // ensemble_logits_buf: [K * B * num_atoms] — all K heads' value logits let na = dqn.config.num_atoms; - let total_actions = dqn.config.num_actions + dqn.config.num_order_types + dqn.config.num_urgency_levels; + // total_actions removed — ensemble buffers use num_atoms, not total_actions. let logits_buf = stream.alloc_zeros::(k * batch_size * na) .map_err(|e| anyhow::anyhow!("Alloc ensemble_logits_buf f32: {e}"))?; - let mean_q_buf = stream.alloc_zeros::(batch_size * total_actions) + // mean_q_buf and var_q_buf: [B * num_atoms] — aggregate kernel outputs + // per-atom mean/variance across K heads (NOT per-action). + let mean_q_buf = stream.alloc_zeros::(batch_size * na) .map_err(|e| anyhow::anyhow!("Alloc ensemble_mean_q_buf: {e}"))?; - let var_q_buf = stream.alloc_zeros::(batch_size * total_actions) + let var_q_buf = stream.alloc_zeros::(batch_size * na) .map_err(|e| anyhow::anyhow!("Alloc ensemble_var_q_buf: {e}"))?; let div_loss_buf = stream.alloc_zeros::(1) .map_err(|e| anyhow::anyhow!("Alloc ensemble_diversity_loss_buf: {e}"))?;