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:
@@ -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
|
||||
|
||||
@@ -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!"
|
||||
|
||||
Reference in New Issue
Block a user