perf: graph all remaining ops — IQN full pipeline + regime_scale in graph_adam

IQN graph now captures the FULL pipeline in one graph:
- decode_actions + fwd/loss + backward + grad_norm + Adam
- trunk gradient (cuBLAS backward into shared weights)
- target EMA (tau from GPU-resident tau_buf, async HtoD before replay)
- IQN→PER loss DtoD copy

IQN EMA kernel changed: float tau → const float* tau_buf (device read).
tau_buf added to GpuIqnHead with async cuMemcpyHtoDAsync per step.
This was the last scalar parameter preventing full graph capture.

regime_scale_td_errors moved into graph_adam submit sequence.
Runs after Adam unflatten, before PER priority update.

Per-step: 7 graph replays + ~9 ungraphed ops
Ungraphed ops (genuinely can't be graphed — batch ptrs change):
  - upload_batch_gpu: 6 DtoD + 2 pad_states (batch-specific pointers)
  - HER relabel: 1-2 kernels (donor from batch next_states)
  - PER priority update: 1 kernel (indices from batch)

Co-Authored-By: Claude Opus 4.6 (1M context) <noreply@anthropic.com>
This commit is contained in:
jgrusewski
2026-04-02 10:06:56 +02:00
parent fb05631997
commit 13dd2e77bf
4 changed files with 71 additions and 42 deletions

View File

@@ -4348,6 +4348,9 @@ impl GpuDqnTrainer {
// ── 7. Unflatten: params_bf16 → individual bf16 weight tensors ─
self.unflatten_online_weights(online_d, online_b)?;
// ── 8. Regime-adaptive PER scaling (element-wise, all pre-allocated) ─
self.regime_scale_td_errors()?;
Ok(())
}

View File

@@ -189,6 +189,7 @@ pub struct GpuIqnHead {
// ── Training state ───────────────────────────────────────────────
adam_step: i32,
t_buf: cudarc::driver::CudaSlice<i32>,
tau_buf: cudarc::driver::CudaSlice<f32>,
/// Monotonic step counter for Philox PRNG seeding (τ sampling).
rng_step: u32,
total_params: usize,
@@ -279,6 +280,8 @@ impl GpuIqnHead {
let t_buf = stream.alloc_zeros::<i32>(1)
.map_err(|e| MLError::ModelError(format!("iqn_t_buf alloc: {e}")))?;
let tau_buf = stream.alloc_zeros::<f32>(1)
.map_err(|e| MLError::ModelError(format!("iqn_tau_buf alloc: {e}")))?;
Ok(Self {
config,
stream,
@@ -312,6 +315,7 @@ impl GpuIqnHead {
total_loss,
adam_step: 0,
t_buf,
tau_buf,
rng_step: 0,
total_params,
})
@@ -591,6 +595,14 @@ impl GpuIqnHead {
/// `target[i] = (1 - tau) * target[i] + tau * online[i]`
pub fn target_ema_update(&mut self, tau: f32) -> Result<(), MLError> {
let n = self.total_params;
unsafe {
cudarc::driver::sys::cuMemcpyHtoDAsync_v2(
self.tau_buf.raw_ptr(),
(&tau as *const f32).cast(),
std::mem::size_of::<f32>(),
self.stream.cu_stream(),
);
}
let blocks = (n + 255) / 256;
let config = LaunchConfig {
grid_dim: (blocks as u32, 1, 1),
@@ -602,6 +614,7 @@ impl GpuIqnHead {
let shared_h1_i32 = self.config.shared_h1 as i32;
let hidden_dim_i32 = self.config.hidden_dim as i32;
let embed_dim_i32 = self.config.embed_dim as i32;
let tau_ptr = self.tau_buf.raw_ptr();
// Safety: target_params and online_params have total_params elements each.
unsafe {
@@ -609,7 +622,7 @@ impl GpuIqnHead {
.launch_builder(&self.ema_kernel)
.arg(&mut self.target_params)
.arg(&self.online_params)
.arg(&tau)
.arg(&tau_ptr)
.arg(&n_i32)
.arg(&shared_h1_i32)
.arg(&hidden_dim_i32)
@@ -639,6 +652,8 @@ impl GpuIqnHead {
}
/// Raw device pointer to IQN trunk gradient — avoids CudaSlice borrow conflicts.
pub fn tau_buf_ptr(&self) -> u64 { self.tau_buf.raw_ptr() }
pub fn d_h_s2_raw_ptr(&self) -> u64 {
self.d_h_s2_buf.raw_ptr()
}

View File

@@ -875,16 +875,16 @@ extern "C" __global__
void iqn_ema_kernel(
__nv_bfloat16* __restrict__ target,
const __nv_bfloat16* __restrict__ online,
float tau,
const float* __restrict__ tau_buf,
int n,
int shared_h1, /* runtime: unused, for consistent interface */
int hidden_dim, /* runtime: unused, for consistent interface */
int embed_dim /* runtime: unused, for consistent interface */
int shared_h1,
int hidden_dim,
int embed_dim
)
{
int i = blockIdx.x * blockDim.x + threadIdx.x;
if (i < n) {
__nv_bfloat16 bf16_tau = bf16(tau);
__nv_bfloat16 bf16_tau = bf16(tau_buf[0]);
target[i] = (bf16_one() - bf16_tau) * target[i] + bf16_tau * online[i];
}
}

View File

@@ -765,14 +765,31 @@ impl FusedTrainingCtx {
// IQN: most of the pipeline is graphed, but trunk gradient + target EMA
// use tau (changes per step) and write to online_dueling (shared state).
// Graph the main IQN forward+loss+backward+adam, keep trunk grad + EMA ungraphed.
// IQN: full pipeline + trunk gradient + EMA + PER DtoD in one graph.
// tau uploaded async before replay (GPU-resident tau_buf).
if self.graph_iqn.is_none() && self.gpu_iqn.is_some() {
// First step: run everything ungraphed
if let Some(ref mut iqn) = self.gpu_iqn {
let _ = iqn.train_iqn_step_gpu(
self.trainer.save_h_s2(), self.trainer.next_states_buf(),
&self.target_dueling, self.trainer.actions_buf(),
self.trainer.rewards_buf(), self.trainer.dones_buf(),
);
let d_ptr = iqn.d_h_s2_raw_ptr();
let _ = self.trainer.apply_iqn_trunk_gradient(d_ptr, &mut self.online_dueling);
let _ = iqn.target_ema_update(0.005); // initial tau
// IQN→PER DtoD
let bs = self.trainer.batch_size();
let n_bytes = bs * std::mem::size_of::<half::bf16>();
unsafe {
let _ = cudarc::driver::result::memcpy_dtod_async(
self.trainer.td_errors_buf().raw_ptr(),
iqn.per_sample_loss().raw_ptr(),
n_bytes, self.stream.cu_stream(),
);
}
}
// Capture the full IQN pipeline as one graph
self.stream.synchronize()
.map_err(|e| anyhow::anyhow!("sync before iqn capture: {e}"))?;
unsafe { self.stream.context().disable_event_tracking(); }
@@ -785,49 +802,47 @@ impl FusedTrainingCtx {
&self.target_dueling, self.trainer.actions_buf(),
self.trainer.rewards_buf(), self.trainer.dones_buf(),
);
let d_ptr = iqn.d_h_s2_raw_ptr();
let _ = self.trainer.apply_iqn_trunk_gradient(d_ptr, &mut self.online_dueling);
let _ = iqn.target_ema_update(0.005);
let bs = self.trainer.batch_size();
let n_bytes = bs * std::mem::size_of::<half::bf16>();
unsafe {
let _ = cudarc::driver::result::memcpy_dtod_async(
self.trainer.td_errors_buf().raw_ptr(),
iqn.per_sample_loss().raw_ptr(),
n_bytes, self.stream.cu_stream(),
);
}
}
if let Ok(Some(graph)) = self.stream.end_capture(
cudarc::driver::sys::CUgraphInstantiate_flags::CUDA_GRAPH_INSTANTIATE_FLAG_AUTO_FREE_ON_LAUNCH,
) {
use crate::cuda_pipeline::gpu_dqn_trainer::SendSyncGraph;
self.graph_iqn = Some(SendSyncGraph(graph));
tracing::info!("Captured IQN CUDA graph (decode + fwd/loss + bwd + adam)");
tracing::info!("Captured IQN CUDA graph (train + trunk grad + EMA + PER DtoD)");
}
}
unsafe { self.stream.context().enable_event_tracking(); }
let _ = self.stream.context().check_err();
} else if let Some(ref graph) = self.graph_iqn {
graph.0.launch().map_err(|e| anyhow::anyhow!("IQN graph replay: {e}"))?;
}
// IQN trunk gradient + target EMA (ungraphed — tau changes per step)
if let Some(ref mut iqn) = self.gpu_iqn {
let d_h_s2_ptr = iqn.d_h_s2_raw_ptr();
self.trainer.apply_iqn_trunk_gradient(
d_h_s2_ptr,
&mut self.online_dueling,
).map_err(|e| anyhow::anyhow!("IQN trunk gradient: {e}"))?;
let dqn = agent.primary_dqn_mut();
let tau = compute_cosine_annealed_tau(
dqn.get_training_steps(),
dqn.config.tau, dqn.config.tau_final, dqn.config.tau_anneal_steps,
);
iqn.target_ema_update(tau as f32)
.map_err(|e| anyhow::anyhow!("IQN EMA update: {e}"))?;
}
// IQN→PER loss DtoD (ungraphed — td_errors_buf address could shift if PER resizes)
if let Some(ref mut iqn) = self.gpu_iqn {
let bs = self.trainer.batch_size();
let n_bytes = bs * std::mem::size_of::<half::bf16>();
let src_ptr = iqn.per_sample_loss().raw_ptr();
let dst_ptr = self.trainer.td_errors_buf().raw_ptr();
unsafe {
cudarc::driver::result::memcpy_dtod_async(
dst_ptr, src_ptr, n_bytes, self.stream.cu_stream()
).map_err(|e| anyhow::anyhow!("IQN→PER loss DtoD: {e}"))?;
// Upload tau before replay (IQN EMA reads from tau_buf)
if let Some(ref iqn) = self.gpu_iqn {
let dqn = agent.primary_dqn_mut();
let tau = compute_cosine_annealed_tau(
dqn.get_training_steps(),
dqn.config.tau, dqn.config.tau_final, dqn.config.tau_anneal_steps,
);
unsafe {
cudarc::driver::sys::cuMemcpyHtoDAsync_v2(
iqn.tau_buf_ptr(),
(&(tau as f32) as *const f32).cast(),
std::mem::size_of::<f32>(),
self.stream.cu_stream(),
);
}
}
graph.0.launch().map_err(|e| anyhow::anyhow!("IQN graph replay: {e}"))?;
}
// ── Step 5b2: Ensemble multi-head diversity loss ──────────────
@@ -880,11 +895,7 @@ impl FusedTrainingCtx {
let fused_result = self.trainer.replay_adam_and_readback()
.map_err(|e| { eprintln!("!!! ADAM REPLAY FAILED: {e}"); anyhow::anyhow!("graph_adam replay: {e}") })?;
// ── Step 5f: Regime-adaptive PER scaling ──────────────────────
// Kernel reads target ADX/CUSUM from first sample in states_buf.
// Zero CPU readback — fully GPU-native.
self.trainer.regime_scale_td_errors()
.map_err(|e| anyhow::anyhow!("Regime PER scaling: {e}"))?;
// Regime PER scaling now captured in graph_adam — no per-step call needed.
// ── Step 6: GPU-native PER priority update ─────────────────────
// td_errors stay on GPU (td_errors_buf). Single CUDA kernel scatter-writes