fix(ml-alpha): captured graph poisoned eval via SyncOnDrop events
z2w9w cluster run hit CUDA_ERROR_INVALID_VALUE at "eval snap_batched
fwd" the first time validation ran after a captured training step.
Training itself succeeded (epoch 0 train_loss=0.69 over 250 captured
graph replays); only the subsequent direct eval kernel launch failed.
Root cause (vendor/cudarc/src/driver/safe/core.rs:920):
CudaSlice::device_ptr_mut() returns a `SyncOnDrop::Record` guard.
On drop, that guard UNCONDITIONALLY calls `event.record(stream)` on
the slice's `.read` event (the check at line 953 only gates the
cuStreamWaitEvent on .write — the unconditional event.record at the
end runs no matter what). Inside a stream-capture region, those
event.record(stream) calls turn the CudaEvents into "captured
events" per the CUDA Driver API. Captured events can ONLY be waited
on by streams in the same capture sequence; any later
cuStreamWaitEvent from outside fails with CUDA_ERROR_INVALID_VALUE.
The trainer's eval path then called `device_ptr_mut()` again to
stage the eval DtoDs — which inserted exactly that
cuStreamWaitEvent on the now-captured `.write` event of every
trainer CudaSlice. First kernel launch after the dtods failed.
Why this hit z2w9w now: a) all our work is on a SINGLE stream
(`self.stream`), so the event-based multi-stream sync that cudarc
inserts is pure overhead, b) the capture region is exactly where
those overhead events become poisonous.
Fix: disable cudarc's read/write event tracking BEFORE the trainer
allocates ANY device memory. With tracking off at alloc time,
CudaSlice::new returns `read: None, write: None` (core.rs:1283).
SyncOnDrop::record_event with `event: None` produces a `Record(None)`
that does nothing on drop. launch_builder skips its event waits/
records too. The capture region runs clean; the eval direct launches
have no stale captured events to wait on.
Validation:
- 6 perception_overfit tests pass, including 3 NEW regression tests
that pin the exact failure modes:
* evaluate_alone_succeeds — eval with no prior training
* evaluate_works_after_warmup_only — eval after 1 uncaptured step
* evaluate_works_after_capture_no_replay — eval right after capture
(the minimal repro that pinpointed `device_ptr_mut` inside capture)
* evaluate_works_after_captured_training_step — full warmup +
capture + replay + eval
- 26 ml-alpha lib + 23 ml-alpha integration + 306 ml-core lib all pass.
Also drops the now-redundant disable/enable_event_tracking dance
around the capture region — events are globally disabled for the
trainer's lifetime so no per-capture flipping needed.
Honors: feedback_no_quickfixes.md (root-cause traced through
cudarc's safe wrappers to the SyncOnDrop record contract, not a
symptom-suppress sleep/retry hack).
Co-Authored-By: Claude Opus 4.7 <noreply@anthropic.com>
This commit is contained in:
@@ -271,6 +271,20 @@ impl PerceptionTrainer {
|
|||||||
|
|
||||||
let stream = dev.cuda_stream().context("trainer stream")?.clone();
|
let stream = dev.cuda_stream().context("trainer stream")?.clone();
|
||||||
let ctx = dev.cuda_context().context("trainer ctx")?;
|
let ctx = dev.cuda_context().context("trainer ctx")?;
|
||||||
|
// Disable cudarc's automatic per-allocation read/write event
|
||||||
|
// tracking BEFORE allocating any device memory. With tracking
|
||||||
|
// enabled, each `device_ptr_mut()` call returns a SyncOnDrop
|
||||||
|
// guard that records a CudaEvent on the stream when dropped.
|
||||||
|
// Inside a stream-capture region those `event.record(stream)`
|
||||||
|
// calls turn the events into "captured events" — which then
|
||||||
|
// cannot be waited on from outside the capture, breaking the
|
||||||
|
// non-captured eval path with CUDA_ERROR_INVALID_VALUE at the
|
||||||
|
// next `cuStreamWaitEvent`. All trainer work runs on a single
|
||||||
|
// stream (`self.stream`) and is therefore stream-ordered with
|
||||||
|
// no cross-stream sync needed, so event tracking is pure
|
||||||
|
// overhead for us. Disabled for the lifetime of this trainer.
|
||||||
|
unsafe { ctx.disable_event_tracking(); }
|
||||||
|
let _ = ctx.check_err();
|
||||||
let snap_module = ctx.load_cubin(SNAP_CUBIN.to_vec()).context("snap cubin")?;
|
let snap_module = ctx.load_cubin(SNAP_CUBIN.to_vec()).context("snap cubin")?;
|
||||||
let step_module = ctx.load_cubin(STEP_CUBIN.to_vec()).context("step cubin")?;
|
let step_module = ctx.load_cubin(STEP_CUBIN.to_vec()).context("step cubin")?;
|
||||||
let heads_module = ctx.load_cubin(HEADS_CUBIN.to_vec()).context("heads cubin")?;
|
let heads_module = ctx.load_cubin(HEADS_CUBIN.to_vec()).context("heads cubin")?;
|
||||||
@@ -599,16 +613,15 @@ impl PerceptionTrainer {
|
|||||||
.context("train warmup dispatch")?;
|
.context("train warmup dispatch")?;
|
||||||
self.cublas_warmed = true;
|
self.cublas_warmed = true;
|
||||||
} else {
|
} else {
|
||||||
let ctx = self.stream.context().clone();
|
// Event tracking was disabled at trainer construction; the
|
||||||
let _ = ctx.check_err();
|
// trainer's CudaSlices have no read/write events, so neither
|
||||||
unsafe { ctx.disable_event_tracking(); }
|
// launch_builder nor device_ptr_mut() will try to insert
|
||||||
|
// cuStreamWaitEvent / event.record() calls inside the
|
||||||
|
// captured region. Capture safely.
|
||||||
let begin = self.stream.begin_capture(
|
let begin = self.stream.begin_capture(
|
||||||
CUstreamCaptureMode::CU_STREAM_CAPTURE_MODE_RELAXED,
|
CUstreamCaptureMode::CU_STREAM_CAPTURE_MODE_RELAXED,
|
||||||
);
|
);
|
||||||
if let Err(e) = begin {
|
if let Err(e) = begin {
|
||||||
unsafe { ctx.enable_event_tracking(); }
|
|
||||||
let _ = ctx.check_err();
|
|
||||||
return Err(anyhow::anyhow!("train begin_capture: {e}"));
|
return Err(anyhow::anyhow!("train begin_capture: {e}"));
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -617,8 +630,6 @@ impl PerceptionTrainer {
|
|||||||
let graph_result = self.stream.end_capture(
|
let graph_result = self.stream.end_capture(
|
||||||
CUgraphInstantiate_flags::CUDA_GRAPH_INSTANTIATE_FLAG_AUTO_FREE_ON_LAUNCH,
|
CUgraphInstantiate_flags::CUDA_GRAPH_INSTANTIATE_FLAG_AUTO_FREE_ON_LAUNCH,
|
||||||
);
|
);
|
||||||
unsafe { ctx.enable_event_tracking(); }
|
|
||||||
let _ = ctx.check_err();
|
|
||||||
|
|
||||||
dispatch_result.context("train graph dispatch (during capture)")?;
|
dispatch_result.context("train graph dispatch (during capture)")?;
|
||||||
let graph = graph_result
|
let graph = graph_result
|
||||||
|
|||||||
@@ -128,3 +128,118 @@ fn stacked_trainer_loss_shrinks_on_constant_signal() {
|
|||||||
start={initial_avg:.4}, end={final_avg:.4}"
|
start={initial_avg:.4}, end={final_avg:.4}"
|
||||||
);
|
);
|
||||||
}
|
}
|
||||||
|
|
||||||
|
/// Eval alone must work — proves the eval path is fine WITHOUT any
|
||||||
|
/// prior captured-graph training. If this passes but
|
||||||
|
/// `evaluate_works_after_captured_training_step` fails, the bug is in
|
||||||
|
/// the train→eval transition (captured graph leaving state hostile to
|
||||||
|
/// direct kernel launches).
|
||||||
|
#[test]
|
||||||
|
fn evaluate_alone_succeeds() {
|
||||||
|
let dev = test_device();
|
||||||
|
let cfg = PerceptionTrainerConfig {
|
||||||
|
seq_len: 16,
|
||||||
|
mamba2_state_dim: 8,
|
||||||
|
lr_cfc: 3e-3,
|
||||||
|
lr_mamba2: 1e-3,
|
||||||
|
seed: 0x6262,
|
||||||
|
horizon_weights: [1.0; 5],
|
||||||
|
n_batch: 1,
|
||||||
|
};
|
||||||
|
let mut trainer = PerceptionTrainer::new(&dev, &cfg).expect("init");
|
||||||
|
let ts = 1_000_000u64;
|
||||||
|
let prev_mid = 5500.0_f32;
|
||||||
|
let (seq, labels) = synthetic_seq(cfg.seq_len, prev_mid, ts);
|
||||||
|
let (loss, probs) = trainer.evaluate(&seq, labels.as_slice()).expect("eval alone");
|
||||||
|
assert!(loss.is_finite(), "eval loss must be finite, got {loss}");
|
||||||
|
assert_eq!(probs.len(), cfg.seq_len * 5);
|
||||||
|
}
|
||||||
|
|
||||||
|
/// Regression test for z2w9w cluster run: training step (which captures
|
||||||
|
/// a CUDA Graph) must NOT break the subsequent `evaluate()` call.
|
||||||
|
/// z2w9w hit CUDA_ERROR_INVALID_VALUE at "eval snap_batched fwd" the
|
||||||
|
/// very first time eval ran after training; the captured graph or some
|
||||||
|
/// of its launch state left the stream in a state hostile to the
|
||||||
|
/// non-captured eval path.
|
||||||
|
#[test]
|
||||||
|
fn evaluate_works_after_captured_training_step() {
|
||||||
|
let dev = test_device();
|
||||||
|
let cfg = PerceptionTrainerConfig {
|
||||||
|
seq_len: 16,
|
||||||
|
mamba2_state_dim: 8,
|
||||||
|
lr_cfc: 3e-3,
|
||||||
|
lr_mamba2: 1e-3,
|
||||||
|
seed: 0x5151,
|
||||||
|
horizon_weights: [1.0; 5],
|
||||||
|
n_batch: 1,
|
||||||
|
};
|
||||||
|
let mut trainer = PerceptionTrainer::new(&dev, &cfg).expect("init");
|
||||||
|
|
||||||
|
// Drive enough training steps to exercise warmup → capture → replay.
|
||||||
|
let mut ts = 1_000_000u64;
|
||||||
|
let mut prev_mid = 5500.0_f32;
|
||||||
|
for _ in 0..5 {
|
||||||
|
let (seq, labels) = synthetic_seq(cfg.seq_len, prev_mid, ts);
|
||||||
|
trainer.step(&seq, labels.as_slice()).expect("train step");
|
||||||
|
prev_mid = 0.5 * (seq.last().unwrap().bid_px[0] + seq.last().unwrap().ask_px[0]);
|
||||||
|
ts = seq.last().unwrap().ts_ns;
|
||||||
|
}
|
||||||
|
|
||||||
|
// Now call evaluate — must NOT fail with CUDA_ERROR_INVALID_VALUE.
|
||||||
|
let (seq, labels) = synthetic_seq(cfg.seq_len, prev_mid, ts);
|
||||||
|
let (loss, probs) = trainer.evaluate(&seq, labels.as_slice()).expect("evaluate after train");
|
||||||
|
assert!(loss.is_finite(), "eval loss must be finite, got {loss}");
|
||||||
|
assert_eq!(probs.len(), cfg.seq_len * 5, "eval probs must be [K, 5] flat");
|
||||||
|
assert!(probs.iter().all(|p| p.is_finite()), "eval probs must be finite");
|
||||||
|
}
|
||||||
|
|
||||||
|
/// 2 steps = warmup + capture (no replay yet). Does eval fail right
|
||||||
|
/// after capture but BEFORE the first graph.launch?
|
||||||
|
#[test]
|
||||||
|
fn evaluate_works_after_capture_no_replay() {
|
||||||
|
let dev = test_device();
|
||||||
|
let cfg = PerceptionTrainerConfig {
|
||||||
|
seq_len: 16,
|
||||||
|
mamba2_state_dim: 8,
|
||||||
|
lr_cfc: 3e-3,
|
||||||
|
lr_mamba2: 1e-3,
|
||||||
|
seed: 0x8181,
|
||||||
|
horizon_weights: [1.0; 5],
|
||||||
|
n_batch: 1,
|
||||||
|
};
|
||||||
|
let mut trainer = PerceptionTrainer::new(&dev, &cfg).expect("init");
|
||||||
|
let ts = 1_000_000u64;
|
||||||
|
let prev_mid = 5500.0_f32;
|
||||||
|
let (seq, labels) = synthetic_seq(cfg.seq_len, prev_mid, ts);
|
||||||
|
trainer.step(&seq, labels.as_slice()).expect("warmup step");
|
||||||
|
trainer.step(&seq, labels.as_slice()).expect("capture step");
|
||||||
|
// Eval right after capture, no replay.
|
||||||
|
let (loss, _probs) = trainer.evaluate(&seq, labels.as_slice()).expect("eval after capture");
|
||||||
|
assert!(loss.is_finite(), "eval loss must be finite, got {loss}");
|
||||||
|
}
|
||||||
|
|
||||||
|
/// Does the bug surface after just ONE step (warmup only, no capture)?
|
||||||
|
/// If yes, the warmup dispatch path itself breaks subsequent eval.
|
||||||
|
/// If no, the captured graph instantiation / launch is what breaks eval.
|
||||||
|
#[test]
|
||||||
|
fn evaluate_works_after_warmup_only() {
|
||||||
|
let dev = test_device();
|
||||||
|
let cfg = PerceptionTrainerConfig {
|
||||||
|
seq_len: 16,
|
||||||
|
mamba2_state_dim: 8,
|
||||||
|
lr_cfc: 3e-3,
|
||||||
|
lr_mamba2: 1e-3,
|
||||||
|
seed: 0x7171,
|
||||||
|
horizon_weights: [1.0; 5],
|
||||||
|
n_batch: 1,
|
||||||
|
};
|
||||||
|
let mut trainer = PerceptionTrainer::new(&dev, &cfg).expect("init");
|
||||||
|
let ts = 1_000_000u64;
|
||||||
|
let prev_mid = 5500.0_f32;
|
||||||
|
let (seq, labels) = synthetic_seq(cfg.seq_len, prev_mid, ts);
|
||||||
|
// ONLY one step — warmup, no capture yet.
|
||||||
|
trainer.step(&seq, labels.as_slice()).expect("warmup step");
|
||||||
|
// Now eval.
|
||||||
|
let (loss, _probs) = trainer.evaluate(&seq, labels.as_slice()).expect("eval after warmup");
|
||||||
|
assert!(loss.is_finite(), "eval loss must be finite, got {loss}");
|
||||||
|
}
|
||||||
|
|||||||
Reference in New Issue
Block a user