- Fix EMA kernel: add tau_buf device field, async HtoD via stable host address, pass pointer not scalar (was ILLEGAL_ADDRESS every step) - Fix HER relabel kernel: revert indirect ptr_buf to direct bf16 pointer - Fix PER update kernel: revert indirect ptr_buf to direct u32/bf16 pointers - Remove IQL per-step DtoH readback (cuStreamSynchronize blocks graph capture) - Permanently disable cudarc event tracking (SyncOnDrop safe for capture) - EventTrackingGuard no longer re-enables tracking on drop - Pre-allocate pass1_event/pass3_event (no cuEventCreate per step) - RawCudaGraph: raw CUDA driver API bypassing cudarc bind_to_thread - graph_aux captures ~30 aux kernel launches (HER+clip+EMA+attn+IQL+IQN+CQL) into single CUDA graph, replayed from step 3+ for zero launch overhead - IQN/GpuDqnTrainer: tau_host stable field for graph-captured HtoD - Remove 3 dead indirect pointer kernels from dqn_utility_kernels.cu - Local RTX 3050: 7.5ms/step steady state (batch=64, 200 steps/epoch) Co-Authored-By: Claude Opus 4.6 (1M context) <noreply@anthropic.com>
17 KiB
Mega-Graph Training Loop Refactor Implementation Plan
For agentic workers: REQUIRED SUB-SKILL: Use superpowers:subagent-driven-development (recommended) or superpowers:executing-plans to implement this plan task-by-task. Steps use checkbox (
- [ ]) syntax for tracking.
Goal: Capture the ENTIRE DQN training step in a single CUDA graph, achieving ~1 graph launch per step with zero ungraphed kernel launches (except PER priority update and vaccine).
Architecture: Move all graph capture logic from GpuDqnTrainer into FusedTrainingCtx which owns all state. Use raw CUDA driver API (cuStreamBeginCapture_v2 / cuStreamEndCapture / cuGraphLaunch) to bypass cudarc's bind_to_thread/check_err that conflicts with graph capture. Pre-allocate ALL CudaEvents before capture to avoid illegal cuEventCreate during capture. Disable cudarc event tracking permanently to prevent SyncOnDrop poisoning.
Tech Stack: Rust 1.85, cudarc 0.17.3 (vendored), CUDA 13.0/12.x driver API, cuBLAS
Critical Bugs in Current Committed Code
Before any refactoring, the committed code has bugs that must be fixed:
- EMA kernel arg mismatch:
ema_kernel.cuexpectsconst float* tau_bufbut Rust passesf32scalar via.arg(&tau)→CUDA_ERROR_ILLEGAL_ADDRESS - IQN EMA kernel arg mismatch: Same issue —
iqn_ema_kernelexpectsconst float* tau_bufbut may still receive scalar - IQN/IQL/Attention Adam kernel arg mismatches: Changed to
const int* t_bufbut Rust may still pass scalari32 - PER update kernel signature changed:
per_update_priorities_kernelexpects indirect pointers but callers pass direct pointers
File Structure
Modified Files
| File | Changes |
|---|---|
crates/ml/src/cuda_pipeline/gpu_dqn_trainer.rs |
Fix kernel arg mismatches, pre-allocate events, make submit methods pub(crate), remove internal graph capture, add sync_all_streams |
crates/ml/src/cuda_pipeline/gpu_attention.rs |
Fix remaining device_ptr → raw_ptr |
crates/ml/src/cuda_pipeline/gpu_iqn_head.rs |
Verify kernel arg compatibility |
crates/ml/src/cuda_pipeline/gpu_iql_trainer.rs |
Verify kernel arg compatibility |
crates/ml/src/trainers/dqn/fused_training.rs |
Add RawCudaGraph, mega-graph capture in FusedTrainingCtx, rewrite run_full_step |
crates/ml/src/cuda_pipeline/ema_kernel.cu |
Revert to scalar tau OR fix Rust caller |
crates/ml/src/cuda_pipeline/iqn_dual_head_kernel.cu |
Verify ema/adam kernel signatures |
crates/ml/src/cuda_pipeline/iql_value_kernel.cu |
Verify adam kernel signature |
crates/ml/src/cuda_pipeline/attention_backward_kernel.cu |
Verify adam kernel signature |
crates/ml/src/cuda_pipeline/per_update_kernel.cu |
Revert to direct pointers OR fix Rust caller |
Task 1: Fix EMA kernel arg mismatch (critical — all tests fail without this)
Files:
- Modify:
crates/ml/src/cuda_pipeline/gpu_dqn_trainer.rs:5443-5454(target_ema_update)
The EMA CUDA kernel was changed to read tau from a device buffer (const float* tau_buf) but the Rust code still passes a scalar f32 via .arg(&tau). The kernel interprets the scalar value as a device pointer → ILLEGAL_ADDRESS crash.
Decision: revert the kernel to accept scalar tau, OR fix the Rust caller.
For the mega-graph, we need the device buffer approach (scalar gets baked at capture time). So: fix the Rust caller to use tau_buf.
- Step 1: Check the committed EMA kernel signature
grep "tau_buf\|float tau\|float\* tau" crates/ml/src/cuda_pipeline/ema_kernel.cu
Expected: const float* __restrict__ tau_buf — kernel reads from device buffer.
- Step 2: Check if tau_buf field exists in GpuDqnTrainer
grep -n "tau_buf" crates/ml/src/cuda_pipeline/gpu_dqn_trainer.rs | head -5
Expected: Field tau_buf: CudaSlice<f32> exists (added in earlier commit).
- Step 3: Fix the EMA kernel launch to pass tau_buf pointer
In gpu_dqn_trainer.rs, find target_ema_update function. Replace:
.arg(&tau)
With:
.arg(&tau_ptr)
Where tau_ptr is obtained from the async HtoD upload that writes tau to self.tau_buf.
Add before the launch:
// Async HtoD for tau (graph-capture compatible)
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 tau_ptr = self.tau_buf.raw_ptr();
And change .arg(&tau) to .arg(&tau_ptr).
- Step 4: Compile and test
SQLX_OFFLINE=true cargo check -p ml
SQLX_OFFLINE=true FOXHUNT_TEST_DATA=test_data/futures-baseline cargo test -p ml --lib -- smoke_tests::training_stability::test_gpu_collector_auto_initializes --ignored
Expected: PASS (if no other kernel mismatches)
- Step 5: Commit
git add crates/ml/src/cuda_pipeline/gpu_dqn_trainer.rs
git commit -m "fix: EMA kernel tau passed as device pointer, not scalar — fixes ILLEGAL_ADDRESS"
Task 2: Fix ALL remaining kernel arg mismatches
Files:
- Modify:
crates/ml/src/cuda_pipeline/gpu_iqn_head.rs(IQN adam + ema) - Modify:
crates/ml/src/cuda_pipeline/gpu_iql_trainer.rs(IQL adam) - Modify:
crates/ml/src/cuda_pipeline/gpu_attention.rs(attention adam) - Modify:
crates/ml/src/cuda_pipeline/per_update_kernel.cu(PER update) - Modify:
crates/ml/src/cuda_pipeline/gpu_dqn_trainer.rs(PER update caller)
For each modified CUDA kernel, verify the Rust caller passes the correct type (device pointer vs scalar). Fix any mismatches.
- Step 1: Audit IQN adam kernel (iqn_dual_head_kernel.cu)
Check if iqn_adam_kernel expects const int* adam_t_buf or int adam_t. Cross-reference with the Rust .arg() call in gpu_iqn_head.rs.
- Step 2: Audit IQN ema kernel (iqn_dual_head_kernel.cu)
Check if iqn_ema_kernel expects const float* tau_buf or float tau. Cross-reference with Rust caller.
- Step 3: Audit IQL adam kernel (iql_value_kernel.cu)
Check if iql_adam_kernel expects const int* t_buf or int t. Cross-reference with Rust caller.
- Step 4: Audit attention adam kernel (attention_backward_kernel.cu)
Check if attn_adam_kernel expects const int* adam_t_buf or int adam_t. Cross-reference with Rust caller.
- Step 5: Audit PER update kernel (per_update_kernel.cu)
Check if per_update_priorities_kernel expects indirect pointers (const unsigned long long*) or direct pointers (const unsigned int*, __nv_bfloat16*). If indirect, the Rust caller must pass batch_ptr_buf offsets. If direct, pass CudaSlice references.
- Step 6: Fix ALL mismatches found
For each mismatch, either: a) Revert the kernel to the original signature (if we don't need graph-capture compatibility yet) b) Fix the Rust caller to pass the correct type
- Step 7: Compile and run full smoke test suite
SQLX_OFFLINE=true cargo check --workspace
SQLX_OFFLINE=true FOXHUNT_TEST_DATA=test_data/futures-baseline cargo test -p ml --lib -- smoke_tests --ignored --nocapture 2>&1 | tail -5
Expected: All smoke tests pass.
- Step 8: Commit
git add -A
git commit -m "fix: all CUDA kernel arg mismatches — device ptr vs scalar compatibility"
Task 3: Verify baseline performance with working code
Files: None (testing only)
After fixing all kernel mismatches, run the smoke test and verify per-step timing. This establishes the baseline before the mega-graph refactor.
- Step 1: Run full smoke test
SQLX_OFFLINE=true FOXHUNT_TEST_DATA=test_data/futures-baseline cargo test -p ml --lib -- smoke_tests --ignored --nocapture 2>&1 | grep "PASSED\|FAILED\|per-step\|step breakdown"
Expected: Tests pass. Per-step timing logged.
- Step 2: Record baseline per-step time
Note the per-step fused time from the test output. This is the number we're trying to beat with the mega-graph.
Task 4: Disable cudarc event tracking permanently
Files:
- Modify:
crates/ml/src/trainers/dqn/fused_training.rs(FusedTrainingCtx::new)
Disable cudarc event tracking before ANY CudaSlice allocation. This ensures all buffers have read=None, write=None, preventing SyncOnDrop from calling cuEventRecord during graph capture.
- Step 1: Add event tracking disable at init
In FusedTrainingCtx::new, BEFORE GpuDqnTrainer::new:
let _ = stream.context().check_err();
unsafe { stream.context().disable_event_tracking(); }
- Step 2: Compile and test
SQLX_OFFLINE=true cargo check -p ml
SQLX_OFFLINE=true FOXHUNT_TEST_DATA=test_data/futures-baseline cargo test -p ml --lib -- smoke_tests::training_stability::test_gpu_collector_auto_initializes --ignored
Expected: PASS (existing 2-graph capture still works with tracking disabled)
- Step 3: Commit
git commit -m "perf: disable cudarc event tracking permanently — SyncOnDrop-safe for graph capture"
Task 5: Pre-allocate multi-stream sync events
Files:
- Modify:
crates/ml/src/cuda_pipeline/gpu_dqn_trainer.rs(struct + constructor + submit_forward_ops)
Replace self.stream.record_event(None) calls in submit_forward_ops with pre-allocated events. cuEventCreate during capture is legal but fragile; pre-allocation is cleaner.
- Step 1: Add event fields to GpuDqnTrainer struct
pass1_event: CudaEvent,
pass3_event: CudaEvent,
- Step 2: Allocate in constructor (BEFORE graph capture)
let pass1_event = stream.record_event(Some(sys::CUevent_flags::CU_EVENT_DISABLE_TIMING))?;
let pass3_event = double_dqn_stream.record_event(Some(sys::CUevent_flags::CU_EVENT_DISABLE_TIMING))?;
- Step 3: Replace record_event calls in submit_forward_ops
Replace:
let pass1_done = self.stream.record_event(None)?;
With:
self.pass1_event.record(&self.stream)?;
And use &self.pass1_event instead of &pass1_done.
Same for pass3_done → self.pass3_event.
- Step 4: Compile and test
Expected: PASS
- Step 5: Commit
git commit -m "perf: pre-allocate multi-stream sync events — no cuEventCreate during capture"
Task 6: Make GpuDqnTrainer submit methods pub(crate)
Files:
- Modify:
crates/ml/src/cuda_pipeline/gpu_dqn_trainer.rs
Make all submit methods accessible from FusedTrainingCtx for the mega-graph capture.
- Step 1: Change visibility
pub(crate) fn submit_forward_ops(...)
pub(crate) fn submit_adam_ops(...)
pub(crate) fn submit_indirect_upload_ops(...)
pub(crate) fn submit_cql_ops(...)
pub(crate) fn submit_c51_clip_ops(...)
pub(crate) fn submit_pruning_mask_ops(...)
pub(crate) fn submit_ema_ops(...)
pub(crate) fn flatten_online_weights(...)
pub(crate) params_initialized: bool
pub(crate) graph_forward: Option<SendSyncGraph>
Also add:
pub(crate) fn sync_all_streams(&self) -> Result<(), MLError> {
self.stream.synchronize()?;
self.double_dqn_stream.synchronize()?;
Ok(())
}
pub(crate) fn adam_step_async(&mut self) { ... }
- Step 2: Compile and test
Expected: PASS (no behavior change)
- Step 3: Commit
Task 7: Implement RawCudaGraph and mega-graph capture
Files:
- Modify:
crates/ml/src/trainers/dqn/fused_training.rs
This is the core refactor. Add RawCudaGraph struct that bypasses cudarc, implement capture_mega_graph on FusedTrainingCtx, and rewrite run_full_step.
- Step 1: Add RawCudaGraph struct
Raw CUDA graph handle that uses cuGraphLaunch directly (no cudarc bind_to_thread/check_err):
struct RawCudaGraph {
exec: cudarc::driver::sys::CUgraphExec,
graph: cudarc::driver::sys::CUgraph,
}
impl RawCudaGraph {
fn launch(&self, stream: cudarc::driver::sys::CUstream) -> Result<()> { ... }
}
impl Drop for RawCudaGraph {
fn drop(&mut self) { cuGraphExecDestroy + cuGraphDestroy }
}
unsafe impl Send for RawCudaGraph {}
unsafe impl Sync for RawCudaGraph {}
- Step 2: Add mega_graph field to FusedTrainingCtx
Replace individual graph fields (graph_attention, graph_iql, graph_iqn, graph_ema, graph_her) with:
mega_graph: Option<RawCudaGraph>,
- Step 3: Implement capture_mega_graph method
Uses raw cuStreamBeginCapture_v2 / cuStreamEndCapture (bypasses cudarc):
fn capture_mega_graph(&mut self) -> Result<()> {
self.trainer.sync_all_streams()?;
let _ = self.stream.context().check_err(); // drain stale errors
// Raw begin_capture
unsafe { sys::cuStreamBeginCapture_v2(cu_stream, MODE) };
// Submit ALL ops in sequence:
self.trainer.submit_forward_ops()?; // includes double_dqn_stream multi-stream
// HER relabel
// Attention fwd + bwd + adam
// IQL value step
// IQN full pipeline
self.trainer.submit_cql_ops()?;
self.trainer.submit_c51_clip_ops()?;
self.trainer.submit_pruning_mask_ops()?;
self.trainer.submit_adam_ops(...)?;
self.trainer.submit_ema_ops(...)?;
self.trainer.regime_scale_td_errors()?;
// Raw end_capture + instantiate
unsafe { sys::cuStreamEndCapture(cu_stream, &mut graph) };
unsafe { sys::cuGraphInstantiateWithFlags(&mut exec, graph, 0) };
self.mega_graph = Some(RawCudaGraph { exec, graph });
}
- Step 4: Rewrite run_full_step
New flow:
- Init weights if needed
- Upload batch pointers + params (async HtoD)
- If mega_graph is None: run all ops ungraphed, then capture_mega_graph
- Else: single mega_graph.launch()
- PER priority update (ungraphed)
- Vaccine (ungraphed, 1/10 steps)
- Bookkeeping
- Step 5: Remove individual graph fields and capture blocks
Delete: graph_attention, graph_iql, graph_iqn, graph_ema, graph_her fields and all their capture/replay code in run_full_step.
- Step 6: Compile and test
SQLX_OFFLINE=true cargo check --workspace
SQLX_OFFLINE=true FOXHUNT_TEST_DATA=test_data/futures-baseline cargo test -p ml --lib -- smoke_tests --ignored --nocapture 2>&1 | tail -10
Expected: All smoke tests pass with mega-graph.
- Step 7: Commit
git commit -m "perf: mega-graph — entire training step in single CUDA graph, 1 launch per step"
Task 8: Remove dead code and cleanup
Files:
-
Modify:
crates/ml/src/cuda_pipeline/gpu_dqn_trainer.rs(remove internal capture) -
Modify:
crates/ml/src/trainers/dqn/fused_training.rs(remove debug logging) -
Step 1: Remove capture_training_graphs from GpuDqnTrainer
The mega-graph capture is now in FusedTrainingCtx. Remove the old method and update any remaining callers (e.g., in train_step_gpu) to not call it.
- Step 2: Remove EventTrackingGuard usage
With event tracking permanently disabled, EventTrackingGuard is a no-op. Remove all instances except the struct definition (may be used elsewhere).
- Step 3: Remove debug eprintln! statements
Remove all eprintln!("MEGA:..."), eprintln!("FWD:..."), eprintln!("ATTN FWD:...") debugging lines.
- Step 4: Compile and run full test suite
SQLX_OFFLINE=true cargo check --workspace
SQLX_OFFLINE=true FOXHUNT_TEST_DATA=test_data/futures-baseline cargo test -p ml --lib -- smoke_tests --ignored
Expected: All pass, zero warnings.
- Step 5: Commit and push
git commit -m "refactor: remove dead graph capture code and debug logging"
git push origin main
Task 9: Deploy to H100 and measure performance
Files: None (deployment only)
- Step 1: Launch H100 training
argo submit -n foxhunt --from workflowtemplate/compile-and-train \
-p commit-sha=$(git rev-parse --short HEAD) \
-p model=dqn -p gpu-pool=ci-training-h100 \
-p hyperopt-trials=0 -p train-epochs=5
- Step 2: Verify per-step timing
Check training logs for Training step breakdown line:
-
Before:
per-step: fused=129.0ms(2-graph approach with syncs) -
Target:
per-step: fused=<10ms(mega-graph, zero ungraphed) -
Step 3: Verify training correctness
Check Sharpe, PF, Q-values, action diversity match expected ranges.
Risk Register
| Risk | Impact | Mitigation |
|---|---|---|
| cudarc launch_builder incompatible with capture | High | Raw cuLaunchKernel as fallback; but research shows launch_builder IS compatible with disabled event tracking |
| double_dqn_stream capture isolation | High | Pre-allocated events + cuStreamWaitEvent automatic join (proven working in current code) |
| Kernel arg type mismatches crash GPU | Critical | Task 1-2 fix ALL mismatches before any refactoring |
| PER update can't be graphed (pointer changes) | Low | Stays outside graph — 1 ungraphed kernel per step (acceptable) |
| cudarc check_err poisons during mega-graph capture | Medium | Drain errors before capture + event tracking disabled = no new errors recorded |