use cubecl_common::bytes::Bytes;
use cubecl_core as cubecl;
use cubecl_core::prelude::*;
use cubecl_core::server::Handle;
use cubecl_cuda::CudaRuntime;
use std::sync::Mutex;
static CAPTURE_LOCK: Mutex<()> = Mutex::new(());
#[cube(launch)]
fn add_one(input: &[f32], output: &mut [f32]) {
if ABSOLUTE_POS < output.len() {
output[ABSOLUTE_POS] = input[ABSOLUTE_POS] + 1.0;
}
}
#[cube(launch)]
fn mul_two(input: &[f32], output: &mut [f32]) {
if ABSOLUTE_POS < output.len() {
output[ABSOLUTE_POS] = input[ABSOLUTE_POS] * 2.0;
}
}
#[test]
fn cuda_graph_capture_replay() {
let _guard = CAPTURE_LOCK.lock().unwrap_or_else(|e| e.into_inner());
let client = CudaRuntime::client(&Default::default());
let n = 4usize;
let input = client.create_from_slice(f32::as_bytes(&[1.0, 2.0, 3.0, 4.0]));
let output = client.empty(n * core::mem::size_of::<f32>());
let launch = |client: &ComputeClient<CudaRuntime>| {
add_one::launch::<CudaRuntime>(
client,
CubeCount::Static(1, 1, 1),
CubeDim::new(client, n),
unsafe { BufferArg::from_raw_parts(input.clone(), n) },
unsafe { BufferArg::from_raw_parts(output.clone(), n) },
);
};
client.graph_prepare().expect("graph_prepare");
launch(&client);
let _ = client.read_one(output.clone()).unwrap();
client.start_capture().expect("start_capture");
launch(&client);
let graph = client.stop_capture().expect("stop_capture");
unsafe { graph.replay() };
let out = client.read_one(output.clone()).unwrap();
assert_eq!(f32::from_bytes(&out), &[2.0, 3.0, 4.0, 5.0]);
unsafe { graph.replay() };
let out = client.read_one(output).unwrap();
assert_eq!(f32::from_bytes(&out), &[2.0, 3.0, 4.0, 5.0]);
}
#[test]
fn cuda_graph_capture_growing_the_pool_is_rejected() {
let _guard = CAPTURE_LOCK.lock().unwrap_or_else(|e| e.into_inner());
let client = CudaRuntime::client(&Default::default());
let n = 4usize;
let input = client.create_from_slice(f32::as_bytes(&[1.0, 2.0, 3.0, 4.0]));
let output = client.empty(n * core::mem::size_of::<f32>());
let launch = |client: &ComputeClient<CudaRuntime>| {
add_one::launch::<CudaRuntime>(
client,
CubeCount::Static(1, 1, 1),
CubeDim::new(client, n),
unsafe { BufferArg::from_raw_parts(input.clone(), n) },
unsafe { BufferArg::from_raw_parts(output.clone(), n) },
);
};
client.graph_prepare().expect("graph_prepare");
launch(&client);
let _ = client.read_one(output.clone()).unwrap();
client.start_capture().expect("start_capture");
launch(&client);
let grown = client.empty(3_145_733);
let rejected = client.stop_capture();
assert!(
rejected.is_err(),
"a capture that grew the pool recorded a memory node and is not relaunchable, so \
stop_capture must reject it rather than return a graph that fails on its second replay"
);
drop(grown);
}
#[test]
fn cuda_graph_input_rewrite() {
let _guard = CAPTURE_LOCK.lock().unwrap_or_else(|e| e.into_inner());
let client = CudaRuntime::client(&Default::default());
let n = 4usize;
let input = client.create_from_slice(f32::as_bytes(&[1.0, 2.0, 3.0, 4.0]));
let output = client.empty(n * core::mem::size_of::<f32>());
let launch = |client: &ComputeClient<CudaRuntime>| {
add_one::launch::<CudaRuntime>(
client,
CubeCount::Static(1, 1, 1),
CubeDim::new(client, n),
unsafe { BufferArg::from_raw_parts(input.clone(), n) },
unsafe { BufferArg::from_raw_parts(output.clone(), n) },
);
};
client.graph_prepare().expect("graph_prepare");
launch(&client);
let _ = client.read_one(output.clone()).unwrap();
client.start_capture().expect("start_capture");
launch(&client);
let graph = client.stop_capture().expect("stop_capture");
unsafe { graph.replay() };
let out = client.read_one(output.clone()).unwrap();
assert_eq!(f32::from_bytes(&out), &[2.0, 3.0, 4.0, 5.0]);
client.write(
&input,
Bytes::from_bytes_vec(f32::as_bytes(&[10.0, 20.0, 30.0, 40.0]).to_vec()),
);
unsafe { graph.replay() };
let out = client.read_one(output).unwrap();
assert_eq!(f32::from_bytes(&out), &[11.0, 21.0, 31.0, 41.0]);
}
#[test]
fn cuda_graph_intermediate_recycling() {
let _guard = CAPTURE_LOCK.lock().unwrap_or_else(|e| e.into_inner());
let client = CudaRuntime::client(&Default::default());
let n = 4usize;
let bytes = n * core::mem::size_of::<f32>();
let input = client.create_from_slice(f32::as_bytes(&[1.0, 2.0, 3.0, 4.0]));
let output = client.empty(bytes);
let run = |client: &ComputeClient<CudaRuntime>, tmp: &Handle| {
add_one::launch::<CudaRuntime>(
client,
CubeCount::Static(1, 1, 1),
CubeDim::new(client, n),
unsafe { BufferArg::from_raw_parts(input.clone(), n) },
unsafe { BufferArg::from_raw_parts(tmp.clone(), n) },
);
mul_two::launch::<CudaRuntime>(
client,
CubeCount::Static(1, 1, 1),
CubeDim::new(client, n),
unsafe { BufferArg::from_raw_parts(tmp.clone(), n) },
unsafe { BufferArg::from_raw_parts(output.clone(), n) },
);
};
client.graph_prepare().expect("graph_prepare");
{
let tmp = client.empty(bytes);
run(&client, &tmp);
let _ = client.read_one(output.clone()).unwrap();
}
client.start_capture().expect("start_capture");
let tmp = client.empty(bytes);
run(&client, &tmp);
let graph = client.stop_capture().expect("stop_capture");
drop(tmp);
let sentinels: Vec<Handle> = (0..8)
.map(|_| client.create_from_slice(f32::as_bytes(&[999.0; 4])))
.collect();
unsafe { graph.replay() };
let out_bytes = client.read_one(output).unwrap();
let out = f32::from_bytes(&out_bytes);
println!("graph output: {out:?} (want [4, 6, 8, 10])");
assert_eq!(out, &[4.0, 6.0, 8.0, 10.0], "graph output corrupted");
let clobbered = sentinels.iter().any(|h| {
let bytes = client.read_one(h.clone()).unwrap();
f32::from_bytes(&bytes) == [2.0, 3.0, 4.0, 5.0]
});
println!("a sentinel buffer was clobbered by replay: {clobbered}");
assert!(
!clobbered,
"replay wrote into a live external buffer that reused the graph's \
intermediate slice — buffer retention failed to pin it"
);
}
#[cube(launch)]
fn add_one_tensor(input: &Tensor<f32>, output: &mut Tensor<f32>) {
if ABSOLUTE_POS < input.shape(0) {
output[ABSOLUTE_POS] = input[ABSOLUTE_POS] + 1.0;
}
}
#[test]
fn cuda_graph_many_launches_dynamic_metadata() {
const N: usize = 64; const PASS_LAUNCHES: usize = 150;
fn run_pass(client: &ComputeClient<CudaRuntime>, a: &Handle, b: &Handle) {
for i in 0..PASS_LAUNCHES {
let (src, dst) = if i % 2 == 0 { (a, b) } else { (b, a) };
add_one_tensor::launch(
client,
CubeCount::Static(1, 1, 1),
CubeDim::new_1d(N as u32),
unsafe { TensorArg::from_raw_parts(src.clone(), [1].into(), [N].into()) },
unsafe { TensorArg::from_raw_parts(dst.clone(), [1].into(), [N].into()) },
);
}
}
fn simulate(start: f32, passes: usize) -> (Vec<f32>, Vec<f32>) {
let (mut a, mut b) = (vec![start; N], vec![start; N]);
for _ in 0..passes {
for i in 0..PASS_LAUNCHES {
if i % 2 == 0 {
for j in 0..N {
b[j] = a[j] + 1.0;
}
} else {
for j in 0..N {
a[j] = b[j] + 1.0;
}
}
}
}
(a, b)
}
let _guard = CAPTURE_LOCK.lock().unwrap_or_else(|e| e.into_inner());
let client = CudaRuntime::client(&Default::default());
let a = client.create_from_slice(f32::as_bytes(&vec![0.0f32; N]));
let b = client.create_from_slice(f32::as_bytes(&vec![0.0f32; N]));
client.graph_prepare().expect("graph_prepare");
run_pass(&client, &a, &b);
let (exp_a, exp_b) = simulate(0.0, 1);
assert_eq!(
f32::from_bytes(&client.read_one(a.clone()).unwrap()),
&exp_a[..]
);
assert_eq!(
f32::from_bytes(&client.read_one(b.clone()).unwrap()),
&exp_b[..]
);
client.start_capture().expect("start_capture");
run_pass(&client, &a, &b);
let graph = client.stop_capture().expect("stop_capture");
unsafe { graph.replay() };
unsafe { graph.replay() };
let (exp_a, exp_b) = simulate(0.0, 3);
assert_eq!(
f32::from_bytes(&client.read_one(a.clone()).unwrap()),
&exp_a[..]
);
assert_eq!(
f32::from_bytes(&client.read_one(b.clone()).unwrap()),
&exp_b[..]
);
let fresh = f32::as_bytes(&[100.0f32; N]).to_vec();
client.write(&a, Bytes::from_bytes_vec(fresh.clone()));
client.write(&b, Bytes::from_bytes_vec(fresh));
unsafe { graph.replay() };
let (exp_a, exp_b) = simulate(100.0, 1);
assert_eq!(
f32::from_bytes(&client.read_one(a.clone()).unwrap()),
&exp_a[..]
);
assert_eq!(
f32::from_bytes(&client.read_one(b.clone()).unwrap()),
&exp_b[..]
);
}