mod dummy;
use crate::dummy::{DummyDevice, DummyElementwiseAddition, test_client};
use cubecl_runtime::server::CubeCount;
use cubecl_runtime::server::KernelArguments;
use cubecl_runtime::{local_tuner, tune::LocalTuner};
use dummy::*;
#[test_log::test]
fn created_resource_is_the_same_when_read() {
let client = test_client(&DummyDevice);
let resource = Vec::from([0, 1, 2]);
let resource_description = client.create_from_slice(&resource);
let obtained_resource = client.read_one(resource_description).unwrap().to_vec();
assert_eq!(resource, obtained_resource)
}
#[test_log::test]
fn empty_allocates_memory() {
let client = test_client(&DummyDevice);
let size = 4;
let resource_description = client.empty(size);
let empty_resource = client.read_one(resource_description).unwrap();
assert_eq!(empty_resource.len(), 4);
}
#[test_log::test]
#[serial_test::parallel]
fn execute_elementwise_addition() {
let client = test_client(&DummyDevice);
let lhs = client.create_from_slice(&[0, 1, 2]);
let rhs = client.create_from_slice(&[4, 4, 4]);
let out = client.empty(3);
client.launch(
Box::new(KernelTask::new(DummyElementwiseAddition)),
CubeCount::Static(1, 1, 1),
KernelArguments::new().with_buffers(vec![
lhs.binding(),
rhs.binding(),
out.clone().binding(),
]),
);
let obtained_resource = client.read_one(out).unwrap().to_vec();
assert_eq!(obtained_resource, Vec::from([4, 5, 6]))
}
#[test_log::test]
#[cfg(feature = "std")]
#[serial_test::serial]
fn autotune_basic_addition_execution() {
static TUNER: LocalTuner<String, String> = local_tuner!("autotune_basic_addition_execution");
let client = test_client(&DummyDevice);
let lhs = client.create_from_slice(&[0, 1, 2]);
let rhs = client.create_from_slice(&[4, 4, 4]);
let out = client.empty(3);
let handles = vec![lhs, rhs, out.clone()];
let test_set = TUNER.init(|| {
let client = test_client(&DummyDevice);
let shapes = vec![vec![1, 3], vec![1, 3], vec![1, 3]];
dummy::addition_set(client, shapes)
});
TUNER.execute(&"test".to_string(), &client, test_set, handles);
let obtained_resource = client.read_one(out).unwrap().to_vec();
assert_eq!(obtained_resource, Vec::from([4, 5, 6]));
}
#[test_log::test]
#[cfg(feature = "std")]
#[serial_test::serial]
fn autotune_basic_multiplication_execution() {
static TUNER: LocalTuner<String, String> =
local_tuner!("autotune_basic_multiplication_execution");
let client = test_client(&DummyDevice);
let lhs = client.create_from_slice(&[0, 1, 2]);
let rhs = client.create_from_slice(&[4, 4, 4]);
let out = client.empty(3);
let handles = vec![lhs, rhs, out.clone()];
let test_set = TUNER.init(|| {
let client = test_client(&DummyDevice);
let shapes = vec![vec![1, 3], vec![1, 3], vec![1, 3]];
dummy::multiplication_set(client, shapes)
});
TUNER.execute(&"test".to_string(), &client, test_set, handles);
let obtained_resource = client.read_one(out).unwrap().to_vec();
assert_eq!(obtained_resource, Vec::from([0, 4, 8]));
}
#[test_log::test]
#[cfg(all(feature = "std", autotune_persistence))]
#[serial_test::serial]
fn autotune_resets_when_the_environment_switches() {
use cubecl_runtime::tune::{TuneCacheResult, Tuner};
let first = tempfile::tempdir().unwrap();
let second = tempfile::tempdir().unwrap();
cubecl_environment::environment::set_root(first.path());
let client = test_client(&DummyDevice);
let shapes = vec![vec![1, 3], vec![1, 3], vec![1, 3]];
let set = dummy::addition_set(test_client(&DummyDevice), shapes);
let handles = vec![
client.create_from_slice(&[0, 1, 2]),
client.create_from_slice(&[4, 4, 4]),
client.empty(3),
];
let key = set.generate_key(&handles);
let tuner: Tuner<String> = Tuner::new("environment-switch", "device0");
tuner.check_tune(
&key,
&handles,
&set,
|| set.compute_checksum(),
&client,
None,
);
assert!(matches!(tuner.fastest(&key), TuneCacheResult::Hit { .. }));
cubecl_environment::environment::set_root(second.path());
assert!(matches!(tuner.fastest(&key), TuneCacheResult::Miss));
tuner.check_tune(
&key,
&handles,
&set,
|| set.compute_checksum(),
&client,
None,
);
assert!(matches!(tuner.fastest(&key), TuneCacheResult::Hit { .. }));
cubecl_environment::environment::set_root(first.path());
assert!(matches!(tuner.fastest(&key), TuneCacheResult::Miss));
let rehydrated = tuner.check_tune(
&key,
&handles,
&set,
|| set.compute_checksum(),
&client,
None,
);
assert!(matches!(rehydrated, TuneCacheResult::Hit { .. }));
}
#[test_log::test]
#[cfg(all(feature = "std", not(target_family = "wasm")))]
#[serial_test::serial]
fn autotune_bounds_short_circuit_accepts_first_within_limit() {
static TUNER: LocalTuner<String, String> = local_tuner!("autotune_bounds_short_circuit");
let client = test_client(&DummyDevice);
let lhs = client.create_from_slice(&[0, 1, 2]);
let rhs = client.create_from_slice(&[4, 4, 4]);
let out = client.empty(3);
let handles = vec![lhs, rhs, out.clone()];
let test_set = TUNER.init(|| {
let client = test_client(&DummyDevice);
let shapes = vec![vec![1, 3], vec![1, 3], vec![1, 3]];
dummy::bounded_addition_set_slow_first(client, shapes, 1.0, 1.0)
});
TUNER.execute(&"test".to_string(), &client, test_set, handles);
let obtained = client.read_one(out).unwrap().to_vec();
assert_eq!(obtained, vec![0, 1, 2]);
}
#[test_log::test]
#[cfg(all(feature = "std", not(target_family = "wasm")))]
#[serial_test::serial]
fn autotune_bounds_unreachable_limit_benchmarks_all() {
static TUNER: LocalTuner<String, String> = local_tuner!("autotune_bounds_unreachable_limit");
let client = test_client(&DummyDevice);
let lhs = client.create_from_slice(&[0, 1, 2]);
let rhs = client.create_from_slice(&[4, 4, 4]);
let out = client.empty(3);
let handles = vec![lhs, rhs, out.clone()];
let test_set = TUNER.init(|| {
let client = test_client(&DummyDevice);
let shapes = vec![vec![1, 3], vec![1, 3], vec![1, 3]];
dummy::bounded_addition_set_slow_first(client, shapes, 1e12, 1.0)
});
TUNER.execute(&"test".to_string(), &client, test_set, handles);
let obtained = client.read_one(out).unwrap().to_vec();
assert_eq!(obtained, vec![4, 5, 6]);
}
#[test_log::test]
#[cfg(all(feature = "std", not(target_family = "wasm")))]
#[serial_test::parallel]
fn autotune_short_circuit_disabled_benchmarks_all() {
static TUNER: LocalTuner<String, String> = local_tuner!("autotune_short_circuit_disabled");
let client = test_client(&DummyDevice);
let lhs = client.create_from_slice(&[0, 1, 2]);
let rhs = client.create_from_slice(&[4, 4, 4]);
let out = client.empty(3);
let handles = vec![lhs, rhs, out.clone()];
let test_set = TUNER.init(|| {
let client = test_client(&DummyDevice);
let shapes = vec![vec![1, 3], vec![1, 3], vec![1, 3]];
dummy::bounded_addition_set_no_short_circuit(client, shapes)
});
TUNER.execute(&"test".to_string(), &client, test_set, handles);
let obtained = client.read_one(out).unwrap().to_vec();
assert_eq!(obtained, vec![4, 5, 6]);
}
#[test_log::test]
#[cfg(feature = "std")]
fn profile_reraises_panic_from_profiled_closure() {
let client = test_client(&DummyDevice);
let reraised = std::panic::catch_unwind(std::panic::AssertUnwindSafe(|| {
client.profile(|| panic!("kernel boom"), "test")
}));
let payload = match reraised {
Ok(_) => panic!("a panic in the profiled closure must surface at the caller"),
Err(payload) => payload,
};
assert_eq!(
payload.downcast_ref::<&str>().copied(),
Some("kernel boom"),
"the re-raised panic must carry the original message"
);
}
#[test_log::test]
#[cfg(feature = "std")]
fn profile_returns_ok_on_success() {
let client = test_client(&DummyDevice);
let (value, _duration) = client
.profile(|| 123u32, "ok")
.expect("a successful profiled closure must return Ok");
assert_eq!(value, 123);
}
#[test_log::test]
#[cfg(feature = "std")]
fn exclusive_stays_recoverable_on_task_panic() {
use cubecl_runtime::server::ServerError;
let client = test_client(&DummyDevice);
let result = client.exclusive(|| panic!("exclusive boom"));
match result {
Err(ServerError::Generic { reason, .. }) => assert!(
reason.contains("exclusive boom"),
"the recoverable error must carry the original message, got: {reason}"
),
Err(other) => panic!("expected a recoverable ServerError::Generic, got: {other}"),
Ok(()) => panic!("expected exclusive to return Err on a task panic, not Ok"),
}
}
#[cfg(feature = "std")]
fn fresh_tune_key_uid() -> String {
std::time::SystemTime::now()
.duration_since(std::time::UNIX_EPOCH)
.unwrap()
.as_nanos()
.to_string()
}
#[test_log::test]
#[cfg(feature = "std")]
#[serial_test::serial]
fn autotune_stops_sampling_a_rejected_candidate() {
use std::sync::Arc;
use std::sync::atomic::{AtomicUsize, Ordering};
static TUNER: LocalTuner<String, String> = local_tuner!("autotune_rejected_candidate");
let client = test_client(&DummyDevice);
let lhs = client.create_from_slice(&[0, 1, 2]);
let rhs = client.create_from_slice(&[4, 4, 4]);
let out = client.empty(3);
let handles = vec![lhs, rhs, out.clone()];
let calls = Arc::new(AtomicUsize::new(0));
let calls_set = calls.clone();
let uid = fresh_tune_key_uid();
let test_set = TUNER.init(move || {
let client = test_client(&DummyDevice);
let shapes = vec![vec![1, 3], vec![1, 3], vec![1, 3]];
dummy::addition_set_with_rejected_candidate(client, shapes, uid.clone(), calls_set.clone())
});
TUNER.execute(&"test".to_string(), &client, test_set, handles);
assert_eq!(calls.load(Ordering::Relaxed), 1);
assert_eq!(client.read_one(out).unwrap().to_vec(), vec![4, 5, 6]);
}
#[test_log::test]
#[cfg(feature = "std")]
#[serial_test::serial]
fn autotune_skips_a_candidate_that_fails_compilation() {
static TUNER: LocalTuner<String, String> = local_tuner!("autotune_failing_compilation");
let client = test_client(&DummyDevice);
let lhs = client.create_from_slice(&[0, 1, 2]);
let rhs = client.create_from_slice(&[4, 4, 4]);
let out = client.empty(3);
let handles = vec![lhs, rhs, out.clone()];
let uid = fresh_tune_key_uid();
let test_set = TUNER.init(move || {
let client = test_client(&DummyDevice);
let shapes = vec![vec![1, 3], vec![1, 3], vec![1, 3]];
dummy::addition_set_with_failing_compilation(client, shapes, uid.clone())
});
TUNER.execute(&"test".to_string(), &client, test_set, handles);
client
.flush()
.expect("the launch failure must not survive the profile it happened in");
assert_eq!(client.read_one(out).unwrap().to_vec(), vec![4, 5, 6]);
let after = client
.exclusive(|| 42)
.expect("the device must keep serving after a candidate failed to compile");
assert_eq!(after, 42);
}
#[test_log::test]
#[cfg(all(feature = "std", not(target_family = "wasm")))]
#[serial_test::serial]
fn autotune_stops_sampling_an_eliminated_candidate() {
use cubecl_runtime::config::{CubeClRuntimeConfig, RuntimeConfig};
use std::sync::Arc;
use std::sync::atomic::{AtomicUsize, Ordering};
let bench = CubeClRuntimeConfig::get().autotune.bench.clone();
if !bench.adaptive {
return;
}
let (min_samples, max_samples) = bench.samples();
static TUNER: LocalTuner<String, String> = local_tuner!("autotune_eliminated_candidate");
let client = test_client(&DummyDevice);
let lhs = client.create_from_slice(&[0, 1, 2]);
let rhs = client.create_from_slice(&[4, 4, 4]);
let out = client.empty(3);
let handles = vec![lhs, rhs, out.clone()];
let fast_calls = Arc::new(AtomicUsize::new(0));
let slow_calls = Arc::new(AtomicUsize::new(0));
let fast_set = fast_calls.clone();
let slow_set = slow_calls.clone();
let uid = fresh_tune_key_uid();
let test_set = TUNER.init(move || {
let client = test_client(&DummyDevice);
let shapes = vec![vec![1, 3], vec![1, 3], vec![1, 3]];
dummy::addition_set_with_slow_candidate(
client,
shapes,
uid.clone(),
fast_set.clone(),
slow_set.clone(),
)
});
TUNER.execute(&"test".to_string(), &client, test_set, handles);
let fast = fast_calls.load(Ordering::Relaxed);
let slow = slow_calls.load(Ordering::Relaxed);
assert!(
slow > min_samples,
"the slow candidate was dropped before it earned it: {slow} calls"
);
assert!(
slow < max_samples + 1,
"the slow candidate was sampled to the ceiling: {slow} calls"
);
assert!(
slow < fast,
"the slow candidate kept pace with the survivors: {slow} vs {fast} calls"
);
assert_eq!(client.read_one(out).unwrap().to_vec(), vec![4, 5, 6]);
}
#[test_log::test]
#[cfg(feature = "std")]
#[serial_test::serial]
fn a_dry_run_drops_an_ordinary_launch() {
use cubecl_runtime::dry_run::DryRun;
let client = test_client(&DummyDevice);
let lhs = client.create_from_slice(&[0, 1, 2]);
let rhs = client.create_from_slice(&[4, 4, 4]);
let out = client.create_from_slice(&[9, 9, 9]);
let add = |out: &cubecl_runtime::server::Handle| {
client.launch(
Box::new(KernelTask::new(DummyElementwiseAddition)),
CubeCount::Static(1, 1, 1),
KernelArguments::new().with_buffers(vec![
lhs.clone().binding(),
rhs.clone().binding(),
out.clone().binding(),
]),
);
};
{
let _dry_run = DryRun::new();
add(&out);
assert_eq!(
client.read_one(out.clone()).unwrap().to_vec(),
Vec::from([9, 9, 9]),
"the launch was compiled and then dropped, so the output is untouched"
);
}
add(&out);
assert_eq!(client.read_one(out).unwrap().to_vec(), Vec::from([4, 5, 6]));
}
#[test_log::test]
#[cfg(feature = "std")]
#[serial_test::serial]
fn a_dry_run_still_autotunes() {
use cubecl_runtime::dry_run::DryRun;
static TUNER: LocalTuner<String, String> = local_tuner!("a_dry_run_still_autotunes");
let client = test_client(&DummyDevice);
let test_set = TUNER.init(|| {
let shapes = vec![vec![1, 3], vec![1, 3], vec![1, 3]];
dummy::addition_set(test_client(&DummyDevice), shapes)
});
let lhs = client.create_from_slice(&[0, 1, 2]);
let rhs = client.create_from_slice(&[4, 4, 4]);
let out = client.empty(3);
{
let _dry_run = DryRun::new();
TUNER.execute(
&"test".to_string(),
&client,
test_set.clone(),
vec![lhs.clone(), rhs.clone(), out.clone()],
);
}
TUNER.execute(
&"test".to_string(),
&client,
test_set,
vec![lhs, rhs, out.clone()],
);
assert_eq!(
client.read_one(out).unwrap().to_vec(),
Vec::from([4, 5, 6]),
"the candidates were measured inside the dry run, so the fast one won"
);
}