fix(rl): remove CudaSlice::clone() from graph capture regions

CudaSlice::clone() allocates device memory, which is forbidden during
CUDA stream capture. Inline the launch_l2_norm and
launch_ema_update_per_step calls to avoid the &mut self borrow conflict
that forced the .clone() workaround.

Co-Authored-By: Claude Opus 4.7 <noreply@anthropic.com>
This commit is contained in:
jgrusewski
2026-05-25 23:24:47 +02:00
parent 02987f0111
commit 45ab5dd066
2 changed files with 221 additions and 38 deletions

View File

@@ -3828,8 +3828,32 @@ impl IntegratedTrainer {
// Streaming kurtosis kernel writes directly to
// RL_TD_KURTOSIS_EMA_INDEX; no separate ema_update_per_step
// needed (the streaming kernel IS the EMA — folds across STEPS).
self.launch_kurtosis(&self.td_per_sample_d.clone(), b_size)
.context("launch_kurtosis(td_per_sample_d)")?;
{
let cfg = LaunchConfig {
grid_dim: (1, 1, 1),
block_dim: (1, 1, 1),
shared_mem_bytes: 0,
};
let b_size_i = b_size as i32;
let out_slot = crate::rl::isv_slots::RL_TD_KURTOSIS_EMA_INDEX as i32;
let mean_slot = crate::rl::isv_slots::RL_TD_KURT_STREAM_MEAN_INDEX as i32;
let m2_slot = crate::rl::isv_slots::RL_TD_KURT_STREAM_M2_INDEX as i32;
let m4_slot = crate::rl::isv_slots::RL_TD_KURT_STREAM_M4_INDEX as i32;
let clamp_slot = crate::rl::isv_slots::RL_TD_KURTOSIS_CLAMP_INDEX as i32;
let mut launch = self.stream.launch_builder(&self.rl_kurtosis_streaming_fn);
launch
.arg(&self.td_per_sample_d)
.arg(&mut self.isv_d)
.arg(&b_size_i)
.arg(&out_slot)
.arg(&mean_slot)
.arg(&m2_slot)
.arg(&m4_slot)
.arg(&clamp_slot);
unsafe {
launch.launch(cfg).context("rl_kurtosis_streaming (td) launch")?;
}
}
// kl_pi EMA — Schulman-style approximation
// `mean(log π_old(a) log π_new(a))` over the sampled action.
@@ -3886,13 +3910,28 @@ impl IntegratedTrainer {
.context("ppo_log_ratio_abs_max_b launch")?;
}
}
self.launch_ema_update_per_step(
crate::rl::isv_slots::RL_KL_PI_EMA_INDEX,
RL_LR_CONTROLLER_ALPHA,
&self.ema_input_scratch_d.clone(),
1,
)
.context("ema_update_per_step(kl_pi)")?;
// Inline ema_update_per_step(kl_pi) — avoids &mut self borrow
// conflict that forced .clone() on ema_input_scratch_d.
{
let cfg = LaunchConfig {
grid_dim: (1, 1, 1),
block_dim: (1, 1, 1),
shared_mem_bytes: std::mem::size_of::<f32>() as u32,
};
let slot_i = crate::rl::isv_slots::RL_KL_PI_EMA_INDEX as i32;
let alpha = RL_LR_CONTROLLER_ALPHA;
let b_size_i = 1i32;
let mut launch = self.stream.launch_builder(&self.ema_update_per_step_fn);
launch
.arg(&self.isv_d)
.arg(&slot_i)
.arg(&alpha)
.arg(&self.ema_input_scratch_d)
.arg(&b_size_i);
unsafe {
launch.launch(cfg).context("ema_update_per_step launch (kl_pi)")?;
}
}
// Per-head gradient-norm EMAs (Q, π, V) — feed the
// signal-modulated LR controller (rl_lr_controller's target
@@ -3900,37 +3939,124 @@ impl IntegratedTrainer {
// ss_*_grad_w_d buffers are persistent struct fields; the
// l2_norm launches here read them after the Adam step has
// already consumed them.
//
// Inlined launch_l2_norm + launch_ema_update_per_step to avoid
// CudaSlice::clone() inside graph capture (allocations forbidden).
let k_dqn_hidden = (N_ACTIONS * Q_N_ATOMS * HIDDEN_DIM) as usize;
self.launch_l2_norm(&self.ss_q_grad_w_d.clone(), k_dqn_hidden)
.context("launch_l2_norm(ss_q_grad_w_d)")?;
self.launch_ema_update_per_step(
crate::rl::isv_slots::RL_Q_GRAD_NORM_EMA_INDEX,
RL_LR_CONTROLLER_ALPHA,
&self.ema_input_scratch_d.clone(),
1,
)
.context("ema_update_per_step(q_grad_norm)")?;
{
let block_dim: u32 = 256;
let cfg = LaunchConfig {
grid_dim: (1, 1, 1),
block_dim: (block_dim, 1, 1),
shared_mem_bytes: (block_dim as usize * std::mem::size_of::<f32>()) as u32,
};
let n_i = k_dqn_hidden as i32;
let mut launch = self.stream.launch_builder(&self.rl_l2_norm_fn);
launch
.arg(&self.ss_q_grad_w_d)
.arg(&mut self.ema_input_scratch_d)
.arg(&n_i);
unsafe {
launch.launch(cfg).context("rl_l2_norm launch (ss_q_grad_w_d)")?;
}
}
{
let cfg = LaunchConfig {
grid_dim: (1, 1, 1),
block_dim: (1, 1, 1),
shared_mem_bytes: std::mem::size_of::<f32>() as u32,
};
let slot_i = crate::rl::isv_slots::RL_Q_GRAD_NORM_EMA_INDEX as i32;
let alpha = RL_LR_CONTROLLER_ALPHA;
let b_size_i = 1i32;
let mut launch = self.stream.launch_builder(&self.ema_update_per_step_fn);
launch
.arg(&self.isv_d)
.arg(&slot_i)
.arg(&alpha)
.arg(&self.ema_input_scratch_d)
.arg(&b_size_i);
unsafe {
launch.launch(cfg).context("ema_update_per_step launch (q_grad_norm)")?;
}
}
let n_actions_hidden = (N_ACTIONS * HIDDEN_DIM) as usize;
self.launch_l2_norm(&self.ss_pi_grad_w_d.clone(), n_actions_hidden)
.context("launch_l2_norm(ss_pi_grad_w_d)")?;
self.launch_ema_update_per_step(
crate::rl::isv_slots::RL_PI_GRAD_NORM_EMA_INDEX,
RL_LR_CONTROLLER_ALPHA,
&self.ema_input_scratch_d.clone(),
1,
)
.context("ema_update_per_step(pi_grad_norm)")?;
{
let block_dim: u32 = 256;
let cfg = LaunchConfig {
grid_dim: (1, 1, 1),
block_dim: (block_dim, 1, 1),
shared_mem_bytes: (block_dim as usize * std::mem::size_of::<f32>()) as u32,
};
let n_i = n_actions_hidden as i32;
let mut launch = self.stream.launch_builder(&self.rl_l2_norm_fn);
launch
.arg(&self.ss_pi_grad_w_d)
.arg(&mut self.ema_input_scratch_d)
.arg(&n_i);
unsafe {
launch.launch(cfg).context("rl_l2_norm launch (ss_pi_grad_w_d)")?;
}
}
{
let cfg = LaunchConfig {
grid_dim: (1, 1, 1),
block_dim: (1, 1, 1),
shared_mem_bytes: std::mem::size_of::<f32>() as u32,
};
let slot_i = crate::rl::isv_slots::RL_PI_GRAD_NORM_EMA_INDEX as i32;
let alpha = RL_LR_CONTROLLER_ALPHA;
let b_size_i = 1i32;
let mut launch = self.stream.launch_builder(&self.ema_update_per_step_fn);
launch
.arg(&self.isv_d)
.arg(&slot_i)
.arg(&alpha)
.arg(&self.ema_input_scratch_d)
.arg(&b_size_i);
unsafe {
launch.launch(cfg).context("ema_update_per_step launch (pi_grad_norm)")?;
}
}
self.launch_l2_norm(&self.ss_v_grad_w_d.clone(), HIDDEN_DIM)
.context("launch_l2_norm(ss_v_grad_w_d)")?;
self.launch_ema_update_per_step(
crate::rl::isv_slots::RL_V_GRAD_NORM_EMA_INDEX,
RL_LR_CONTROLLER_ALPHA,
&self.ema_input_scratch_d.clone(),
1,
)
.context("ema_update_per_step(v_grad_norm)")?;
{
let block_dim: u32 = 256;
let cfg = LaunchConfig {
grid_dim: (1, 1, 1),
block_dim: (block_dim, 1, 1),
shared_mem_bytes: (block_dim as usize * std::mem::size_of::<f32>()) as u32,
};
let n_i = HIDDEN_DIM as i32;
let mut launch = self.stream.launch_builder(&self.rl_l2_norm_fn);
launch
.arg(&self.ss_v_grad_w_d)
.arg(&mut self.ema_input_scratch_d)
.arg(&n_i);
unsafe {
launch.launch(cfg).context("rl_l2_norm launch (ss_v_grad_w_d)")?;
}
}
{
let cfg = LaunchConfig {
grid_dim: (1, 1, 1),
block_dim: (1, 1, 1),
shared_mem_bytes: std::mem::size_of::<f32>() as u32,
};
let slot_i = crate::rl::isv_slots::RL_V_GRAD_NORM_EMA_INDEX as i32;
let alpha = RL_LR_CONTROLLER_ALPHA;
let b_size_i = 1i32;
let mut launch = self.stream.launch_builder(&self.ema_update_per_step_fn);
launch
.arg(&self.isv_d)
.arg(&slot_i)
.arg(&alpha)
.arg(&self.ema_input_scratch_d)
.arg(&b_size_i);
unsafe {
launch.launch(cfg).context("ema_update_per_step launch (v_grad_norm)")?;
}
}
if capturing_training {
let graph = self.stream
@@ -5556,8 +5682,34 @@ impl IntegratedTrainer {
// writes directly to ISV[421] (folds across STEPS — works at
// b_size=1 where per-batch variance is undefined). No
// ema_update_per_step needed downstream.
self.launch_var_over_abs_mean(&self.advantages_d.clone(), b_size)
.context("launch_var_over_abs_mean(advantages_d)")?;
{
let cfg = LaunchConfig {
grid_dim: (1, 1, 1),
block_dim: (1, 1, 1),
shared_mem_bytes: 0,
};
let b_size_i = b_size as i32;
let out_slot = crate::rl::isv_slots::RL_ADVANTAGE_VAR_RATIO_EMA_INDEX as i32;
let mean_slot = crate::rl::isv_slots::RL_ADV_VAR_STREAM_MEAN_INDEX as i32;
let m2_slot = crate::rl::isv_slots::RL_ADV_VAR_STREAM_M2_INDEX as i32;
let clamp_slot = crate::rl::isv_slots::RL_ADV_VAR_RATIO_CLAMP_INDEX as i32;
let mut launch = self.stream.launch_builder(
&self.rl_var_over_abs_mean_streaming_fn,
);
launch
.arg(&self.advantages_d)
.arg(&mut self.isv_d)
.arg(&b_size_i)
.arg(&out_slot)
.arg(&mean_slot)
.arg(&m2_slot)
.arg(&clamp_slot);
unsafe {
launch
.launch(cfg)
.context("rl_var_over_abs_mean_streaming (advantages) launch")?;
}
}
if capturing_reward {
let graph = self

View File

@@ -283,11 +283,42 @@ check_no_raw_htod_dtoh() {
fi
}
# Diff-aware guard: CudaSlice::clone() allocates device memory. Inside
# CUDA graph capture regions this causes CUDA_ERROR_STREAM_CAPTURE_UNSUPPORTED.
check_no_cudaslice_clone() {
local staged
staged=$(git diff --cached --name-only --diff-filter=ACM | grep '\.rs$' || true)
if [ -z "$staged" ]; then return 0; fi
local bad=""
while IFS= read -r f; do
[ -z "$f" ] && continue
local hits
hits=$(git diff --cached -U0 -- "$f" 2>/dev/null \
| grep -E '^\+[^+]' \
| grep -E '\._d\.clone\(\)|isv_d\.clone\(\)|scratch_d\.clone\(\)' \
| grep -v '// clone-ok:' \
|| true)
if [ -n "$hits" ]; then
bad+="${bad:+$'\n'}$f:"$'\n'"$hits"
fi
done <<< "$staged"
if [ -n "$bad" ]; then
echo "❌ CudaSlice::clone() allocates device memory — forbidden in graph capture:"
echo "$bad" | sed 's/^/ /'
echo " Fix: inline the kernel launch or use a free function to avoid &mut self borrow."
echo " Suppress with // clone-ok: <reason>"
return 1
fi
}
check_audit_doc_updates || exit 1
check_no_todo_fixme || exit 1
check_no_isv_migrations || exit 1
check_no_dtod_via_pinned || exit 1
check_no_raw_htod_dtoh || exit 1
check_no_cudaslice_clone || exit 1
check_sp18_consumer_audit || exit 1
echo "✅ All pre-commit checks passed!"