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:
@@ -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(())
|
||||
}
|
||||
|
||||
|
||||
@@ -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()
|
||||
}
|
||||
|
||||
@@ -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];
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user