diff --git a/crates/ml/src/trainers/dqn/trainer/training_loop.rs b/crates/ml/src/trainers/dqn/trainer/training_loop.rs index 824524b49..8d00f147c 100644 --- a/crates/ml/src/trainers/dqn/trainer/training_loop.rs +++ b/crates/ml/src/trainers/dqn/trainer/training_loop.rs @@ -1347,23 +1347,11 @@ impl DQNTrainer { self.gradient_logging_step as u64 > ws }; - // Batch pre-sampling: sample batches in chunks under one READ lock. - // usize::MAX caused 488 batches × 8192 samples = 4M PER traversals - // in one go from a 24.7M sum tree — minutes of CPU work on H100. - const PREFETCH_K: usize = 16; let mut sample_total_us = 0_u64; let mut fused_total_us = 0_u64; let mut guard_total_us = 0_u64; - // Task 15: capture regime-biased sampling params before the borrow - let regime_decay = self.regime_replay_decay_override; - let regime_id = self.fold_dominant_regime; - let use_regime_bias = regime_decay < 1.0; - // Wait for experience collection GPU kernels to complete before sampling. - // The event was recorded at the end of collect_gpu_experiences — all CPU-side - // setup above (can_train, ensure_fused_ctx, guard init, variable capture) - // overlaps with the tail of experience collection on the GPU. if let Some(event) = self.experience_done_event.take() { if !event.is_complete() { event.synchronize() @@ -1371,81 +1359,45 @@ impl DQNTrainer { } } - eprintln!("[DEBUG] run_training_steps: num_steps={}, batch_size={}, prefetch_k={}", num_training_steps, batch_size, PREFETCH_K); - for chunk_start in (0..num_training_steps).step_by(PREFETCH_K) { - let chunk_end = (chunk_start + PREFETCH_K).min(num_training_steps); + eprintln!("[DEBUG] run_training_steps: num_steps={}, batch_size={}", num_training_steps, batch_size); + // Sample-train loop: one batch at a time, zero prefetch. + // GPU PER sampling (seg_tree_sample kernel) + GPU training (graph replay) + // interleave naturally. No CPU-side Vec accumulation. + for _step in 0..num_training_steps { + // Sample one batch (GPU-native PER for GpuPrioritized) let sample_start = std::time::Instant::now(); - let (batches, vaccine_batches) = { + let (batch, vaccine_batch) = { let agent = self.agent.read().await; let buffer = agent.memory(); - let can = buffer.can_sample(self.current_batch_size); - let buf_len = buffer.len(); - tracing::debug!( - can_sample = can, - buffer_len = buf_len, - batch_size = self.current_batch_size, - chunk = chunk_start, - regime_biased = use_regime_bias, - "Pre-sampling replay buffer" - ); - let mut b = Vec::with_capacity(chunk_end - chunk_start); - // #32 Gradient Vaccine: always sample a second batch for validation - let mut vaccine_batches = Vec::with_capacity(chunk_end - chunk_start); - for _ in chunk_start..chunk_end { - b.push(buffer.can_sample(self.current_batch_size).then(|| { - // Task 15: Use regime-biased sampling when adaptive decay < 1.0 - if use_regime_bias { - buffer - .sample_regime_biased(self.current_batch_size, regime_id, regime_decay) - .map_err(|e| anyhow::anyhow!("Regime-biased pre-sample: {e}")) - } else { - buffer - .sample(self.current_batch_size) - .map_err(|e| anyhow::anyhow!("Pre-sample: {e}")) - } - }).transpose()?); - // Sample separate vaccine batch (different indices from PER) - // Vaccine uses standard PER (unbiased) for gradient vaccine diversity - if buffer.can_sample(self.current_batch_size) { - vaccine_batches.push(buffer.sample(self.current_batch_size) - .map_err(|e| anyhow::anyhow!("Vaccine sample: {e}")).ok()); - } else { - vaccine_batches.push(None); - } + if !buffer.can_sample(self.current_batch_size) { + break; } - (b, vaccine_batches) + let b = buffer.sample(self.current_batch_size) + .map_err(|e| anyhow::anyhow!("PER sample: {e}"))?; + let vb = buffer.sample(self.current_batch_size).ok(); + (b, vb) }; sample_total_us += sample_start.elapsed().as_micros() as u64; - // GPU train steps (single WRITE lock) + // Train one step { let mut agent = self.agent.write().await; - let accum_steps = self.hyperparams.gradient_accumulation_steps; - if accum_steps <= 1 { - let mut vaccine_iter = vaccine_batches.into_iter(); - for explicit_batch in batches { - let fused_start = std::time::Instant::now(); - let _gpu_result = if let Some(ref mut fused) = self.fused_ctx { - let batch_data = explicit_batch.as_ref().ok_or_else(|| { - anyhow::anyhow!("No batch data for fused training step") - })?; - // #32 Set vaccine batch for this step (consumed by run_full_step) - if let Some(Some(vb)) = vaccine_iter.next() { - fused.pending_vaccine_batch = vb.gpu_batch; - } - let result = fused.run_full_step(batch_data, &mut *agent, &self.device) - .map_err(|e| { eprintln!("!!! FUSED STEP ERROR: {:#}", e); e }) - .context("Fused CUDA training step failed")?; - if train_step_count % 50 == 0 { - } - result - } else { - unreachable!("Fused CUDA training is the only production path — fused_ctx must be initialized") - }; + let fused_start = std::time::Instant::now(); + let _gpu_result = if let Some(ref mut fused) = self.fused_ctx { + if let Some(vb) = vaccine_batch { + fused.pending_vaccine_batch = vb.gpu_batch; + } + let result = fused.run_full_step(&batch, &mut *agent, &self.device) + .map_err(|e| { eprintln!("!!! FUSED STEP ERROR: {:#}", e); e }) + .context("Fused CUDA training step failed")?; + result + } else { + unreachable!("Fused CUDA training is the only production path") + }; - fused_total_us += fused_start.elapsed().as_micros() as u64; + fused_total_us += fused_start.elapsed().as_micros() as u64; let guard_start = std::time::Instant::now(); if let Some(ref mut guard) = self.training_guard { @@ -1480,22 +1432,20 @@ impl DQNTrainer { })?; } } - guard_total_us += guard_start.elapsed().as_micros() as u64; + guard_total_us += guard_start.elapsed().as_micros() as u64; - // Q-value stats: reduce from training batch every 50 steps - if train_step_count % 50 == 0 { - if let Some(ref mut fused) = self.fused_ctx { - if let Ok(stats) = fused.reduce_current_q_stats() { - self.cached_avg_q = stats.q_mean as f64; - if stats.q_min < self.epoch_q_min { self.epoch_q_min = stats.q_min; } - if stats.q_max > self.epoch_q_max { self.epoch_q_max = stats.q_max; } - } - } + // Q-value stats: reduce from training batch every 50 steps + if train_step_count % 50 == 0 { + if let Some(ref mut fused) = self.fused_ctx { + if let Ok(stats) = fused.reduce_current_q_stats() { + self.cached_avg_q = stats.q_mean as f64; + if stats.q_min < self.epoch_q_min { self.epoch_q_min = stats.q_min; } + if stats.q_max > self.epoch_q_max { self.epoch_q_max = stats.q_max; } } - train_step_count += 1; - self.gradient_logging_step += 1; } } + train_step_count += 1; + self.gradient_logging_step += 1; } }