mod dummy;
use crate::dummy::{DummyDevice, DummyElementwiseAddition, test_client};
use cubecl_common::bytes::Bytes;
use cubecl_common::device::{DeviceId, ServiceId};
use cubecl_environment::stream::StreamId;
use cubecl_ir::{ElemType, UIntKind};
use cubecl_server::client::Client;
use cubecl_server::server::{
CubeCount, Handle, IoError, KernelArguments, ReduceOperation, ServerError,
};
use cubecl_server::{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);
}
fn handle_of(service: ServiceId) -> Handle {
Handle::new(service, StreamId::current(), 8)
}
#[test_log::test]
fn a_handle_of_this_client_passes_the_check() {
let client = test_client(&DummyDevice);
let handle = client.empty(8);
assert!(client.check(&[handle]).is_ok());
}
#[test_log::test]
fn a_handle_whose_reservation_failed_is_refused() {
let client = test_client(&DummyDevice);
let served = client.empty(8);
let refused = client.empty(1024 * 1024 * 1024);
assert!(client.check([&served]).is_ok());
let Err(ServerError::Several { errors, .. }) = client.check([&served, &refused]) else {
panic!("the refused buffer passed the check");
};
assert!(matches!(
errors[..],
[ServerError::Io(IoError::NotFound { .. })]
));
assert!(client.read_one(refused).is_err());
}
#[test_log::test]
fn a_handle_from_another_device_of_the_same_runtime_is_refused() {
let client = test_client(&DummyDevice);
let other_device = DeviceId::new(0, 1);
let handle = handle_of(ServiceId::of::<DummyServer>(other_device));
let err = client.check(&[handle]).unwrap_err();
assert!(
matches!(err, ServerError::ForeignHandle { .. }),
"expected the handle to be refused, got: {err}"
);
}
#[test_log::test]
fn a_handle_from_another_runtime_on_the_same_device_id_is_refused() {
let client = test_client(&DummyDevice);
let same_device = client.service_id().device;
let handle = handle_of(ServiceId::of::<()>(same_device));
let err = client.check(&[handle]).unwrap_err();
assert!(
matches!(err, ServerError::ForeignHandle { .. }),
"expected the handle to be refused, got: {err}"
);
}
#[test_log::test]
#[should_panic(expected = "was used on")]
fn writing_through_a_foreign_handle_panics() {
let client = test_client(&DummyDevice);
let handle = handle_of(ServiceId::of::<()>(DeviceId::new(0, 0)));
client.write(&handle, Bytes::from_bytes_vec(vec![0; 8]));
}
#[test_log::test]
#[should_panic(expected = "was used on")]
fn transferring_a_foreign_handle_to_another_client_panics() {
let mut client = test_client(&DummyDevice);
let destination = client.clone();
let handle = handle_of(ServiceId::of::<()>(DeviceId::new(0, 0)));
client.to_client(handle, &destination, ElemType::UInt(UIntKind::U8));
}
#[test_log::test]
fn a_transfer_between_devices_of_the_same_runtime_round_trips() {
let mut source = test_client(&DummyDevice);
let destination = Client::load::<DummyServer>(DeviceId::new(0, 1));
let bytes = [1u8, 2, 3, 4];
let handle = source.create_from_slice(&bytes);
let transferred = source.to_client(handle, &destination, ElemType::UInt(UIntKind::U8));
assert_eq!(transferred.service, destination.service_id());
assert_eq!(destination.read_one(transferred).unwrap().to_vec(), bytes);
}
#[test_log::test]
#[should_panic(expected = "no transport between its devices")]
fn a_transfer_without_a_device_transport_panics_on_the_caller() {
let mut source = test_client(&DummyDevice);
let destination = Client::load::<DummyServer>(DeviceId::new(0, 1));
let handle = source.create_from_slice(&[1u8, 2, 3, 4]);
assert!(!source.has_device_transport());
let descriptor = handle.copy_descriptor([4].into(), [1].into(), 1);
source.to_client_tensor(descriptor, &destination, ElemType::UInt(UIntKind::U8));
}
#[test_log::test]
#[should_panic(expected = "no transport between its devices")]
fn an_all_reduce_without_a_device_transport_panics_on_the_caller() {
let mut client = test_client(&DummyDevice);
let handle = client.create_from_slice(&[1u8, 2, 3, 4]);
let device_ids = vec![DeviceId::new(0, 0), DeviceId::new(0, 1)];
client.all_reduce(
handle.clone(),
handle,
ElemType::UInt(UIntKind::U8),
device_ids,
ReduceOperation::Sum,
);
}
#[test_log::test]
fn a_sync_collective_without_a_device_transport_waits_for_nothing() {
let client = test_client(&DummyDevice);
client.sync_collective();
}
#[test_log::test]
fn a_transfer_across_runtimes_goes_through_the_host() {
let mut source = test_client(&DummyDevice);
let destination = Client::load::<DummyServer<Other>>(DeviceId::new(0, 0));
let bytes = [5u8, 6, 7, 8];
let handle = source.create_from_slice(&bytes);
let transferred = source.to_client(handle, &destination, ElemType::UInt(UIntKind::U8));
assert_eq!(transferred.service, destination.service_id());
assert_eq!(destination.read_one(transferred).unwrap().to_vec(), bytes);
}
#[test_log::test]
fn asking_for_the_resource_of_another_server_type_is_refused() {
let client = test_client(&DummyDevice);
let handle = client.create_from_slice(&[1u8, 2, 3, 4]);
let err = client
.get_resource::<DummyServer<Other>>(handle.clone())
.unwrap_err();
assert!(
matches!(err, ServerError::ServiceMismatch { .. }),
"expected the server type to be refused, got: {err}"
);
assert!(client.get_resource::<DummyServer>(handle.clone()).is_ok());
assert_eq!(client.read_one(handle).unwrap().to_vec(), [1, 2, 3, 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]
#[serial_test::serial]
fn a_refused_profile_degrades_to_an_untimed_launch() {
use cubecl_server::logging::{
Duration, LaunchObservation, LaunchObserver, TimingMethod, TimingRequest,
};
#[derive(Default)]
struct WantsTiming {
launched: std::sync::Mutex<Vec<&'static str>>,
timed: std::sync::Mutex<Vec<&'static str>>,
}
impl LaunchObserver for WantsTiming {
fn launched(&self, kernel: &'static str) {
self.launched.lock().unwrap().push(kernel);
}
fn timing(&self) -> TimingRequest {
TimingRequest::Resolved
}
fn timed(&self, kernel: &'static str, _duration: Duration, _method: TimingMethod) {
self.timed.lock().unwrap().push(kernel);
}
}
let observer = std::sync::Arc::new(WantsTiming::default());
let watching = LaunchObservation::new(observer.clone());
REFUSE_PROFILES.store(true, core::sync::atomic::Ordering::Relaxed);
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();
REFUSE_PROFILES.store(false, core::sync::atomic::Ordering::Relaxed);
drop(watching);
assert_eq!(
obtained_resource,
Vec::from([4, 5, 6]),
"the refused profile must not cost the launch"
);
assert_eq!(
observer.launched.lock().unwrap().len(),
1,
"the launch is still reported"
);
assert!(
observer.timed.lock().unwrap().is_empty(),
"nothing was measured, so nothing is reported as measured"
);
}
#[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(&"test".to_string(), || {
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(&"test".to_string(), || {
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", persistence))]
#[serial_test::serial]
fn autotune_resets_when_the_environment_switches() {
use cubecl_server::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 { .. }));
}
#[cfg(all(feature = "std", persistence))]
fn rooted_at(root: &std::path::Path) {
use cubecl_server::config::RuntimeConfig;
let _ = cubecl_server::config::CubeClRuntimeConfig::get();
cubecl_environment::environment::set_root(root);
}
#[cfg(all(feature = "std", persistence))]
fn recording_at(level: cubecl_environment::records::RecordLevel) {
use cubecl_environment::records::{self, RecordsConfig};
records::configure(RecordsConfig {
level,
..Default::default()
});
}
#[test_log::test]
#[cfg(all(feature = "std", persistence))]
#[serial_test::serial]
fn a_tune_is_recorded_in_order_with_its_walls() {
use cubecl_environment::persistence::Database;
use cubecl_environment::records::{RecordLevel, Records};
use cubecl_server::tune::{TuneCacheResult, TuneRecord, Tuner};
let root = tempfile::tempdir().unwrap();
rooted_at(root.path());
recording_at(RecordLevel::Basic);
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("recorded", "device0");
let answer = tuner.check_tune(
&key,
&handles,
&set,
|| set.compute_checksum(),
&client,
None,
);
let TuneCacheResult::Hit { fastest_index } = answer else {
panic!("the tune answers inline: {answer:?}");
};
let database = Database::open_active().unwrap();
let tunes = Records::new(&database).read::<TuneRecord<String>>();
assert_eq!(tunes.len(), 1);
let tune = &tunes[0].record;
assert_eq!(tune.entry.key, key);
assert_eq!(tune.entry.checksum, set.compute_checksum());
assert_eq!(tune.winner, fastest_index);
assert!(tune.table.starts_with("autotune/") && tune.table.ends_with("device0/recorded"));
let names: Vec<&str> = tune
.trials
.iter()
.map(|trial| trial.name.as_str())
.collect();
assert_eq!(
names,
vec!["add", "add_slow_wrong"],
"in registration order"
);
assert!(tune.trials.iter().all(|trial| !trial.wall.is_zero()));
assert!(
tune.trials
.iter()
.map(|trial| trial.wall)
.sum::<core::time::Duration>()
<= tune.wall
);
assert_eq!(tune.short_circuit, None);
assert!(!tune.dry_run);
assert!(tune.stored, "the table took the answer");
let sessions = Records::new(&database).sessions();
assert_eq!(sessions.len(), 1);
assert_eq!(tunes[0].stamp.session, sessions[0].id);
}
#[test_log::test]
#[cfg(all(feature = "std", persistence, not(target_family = "wasm")))]
#[serial_test::serial]
fn a_short_circuited_tune_records_where_it_stopped() {
use cubecl_environment::persistence::Database;
use cubecl_environment::records::{RecordLevel, Records};
use cubecl_server::tune::{TuneRecord, Tuner};
let root = tempfile::tempdir().unwrap();
rooted_at(root.path());
recording_at(RecordLevel::Basic);
let client = test_client(&DummyDevice);
let shapes = vec![vec![1, 3], vec![1, 3], vec![1, 3]];
let set = dummy::bounded_addition_set_slow_first(test_client(&DummyDevice), shapes, 1.0, 1.0);
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("short-circuited", "device0");
tuner.check_tune(
&key,
&handles,
&set,
|| set.compute_checksum(),
&client,
None,
);
let database = Database::open_active().unwrap();
let tunes = Records::new(&database).read::<TuneRecord<String>>();
let tune = &tunes.last().unwrap().record;
assert_eq!(tune.short_circuit.as_deref(), Some("add_slow_wrong"));
assert!(
tune.trials
.iter()
.any(|trial| trial.name == "add_slow_wrong")
);
}
#[test_log::test]
#[cfg(all(feature = "std", persistence))]
#[serial_test::serial]
fn nothing_is_recorded_when_records_are_off() {
use cubecl_environment::persistence::Database;
use cubecl_environment::records::{RecordLevel, Records};
use cubecl_server::tune::{TuneCacheResult, TuneRecord, Tuner};
let root = tempfile::tempdir().unwrap();
rooted_at(root.path());
recording_at(RecordLevel::Off);
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("unrecorded", "device0");
let answer = tuner.check_tune(
&key,
&handles,
&set,
|| set.compute_checksum(),
&client,
None,
);
recording_at(RecordLevel::Basic);
assert!(matches!(answer, TuneCacheResult::Hit { .. }));
let database = Database::open_active().unwrap();
assert!(
Records::new(&database)
.read::<TuneRecord<String>>()
.is_empty()
);
}
#[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(&"test".to_string(), || {
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(&"test".to_string(), || {
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(&"test".to_string(), || {
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(all(feature = "std", not(target_family = "wasm")))]
#[serial_test::parallel]
fn autotune_evicts_before_every_measured_sample() {
use cubecl_runtime::config::{CubeClRuntimeConfig, RuntimeConfig};
use std::sync::Arc;
use std::sync::atomic::{AtomicUsize, Ordering};
static TUNER: LocalTuner<String, String> = local_tuner!("autotune_eviction");
let candidates = 2;
let warmups = if CubeClRuntimeConfig::get().autotune.bench.adaptive {
1
} else {
3
};
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 evictions = Arc::new(AtomicUsize::new(0));
let misdirected = Arc::new(AtomicUsize::new(0));
let calls_set = calls.clone();
let evictions_set = evictions.clone();
let misdirected_set = misdirected.clone();
let uid = fresh_tune_key_uid();
let test_set = TUNER.init(&"test".to_string(), move || {
let client = test_client(&DummyDevice);
let shapes = vec![vec![1, 3], vec![1, 3], vec![1, 3]];
dummy::addition_set_with_eviction(
client,
shapes,
uid.clone(),
calls_set.clone(),
evictions_set.clone(),
misdirected_set.clone(),
)
});
TUNER.execute(&"test".to_string(), &client, test_set, handles);
assert_eq!(client.read_one(out).unwrap().to_vec(), vec![4, 5, 6]);
let unmeasured = candidates * warmups + 1;
let calls = calls.load(Ordering::Relaxed);
let evictions = evictions.load(Ordering::Relaxed);
let misdirected = misdirected.load(Ordering::Relaxed);
assert_eq!(
misdirected, 0,
"{misdirected} evictions ran on the generated inputs rather than the reference ones"
);
assert!(
calls > unmeasured,
"the candidates were launched {calls} times, no more than the {unmeasured} unmeasured ones"
);
assert_eq!(
evictions,
calls - unmeasured,
"{evictions} evictions for {calls} launches, {unmeasured} of them unmeasured"
);
}
#[test_log::test]
#[cfg(feature = "std")]
#[serial_test::parallel]
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")]
#[serial_test::parallel]
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_server::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(&"test".to_string(), 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(&"test".to_string(), 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(feature = "std")]
#[serial_test::serial]
fn autotune_survives_a_failing_candidate_ahead_of_the_winner() {
static TUNER: LocalTuner<String, String> = local_tuner!("autotune_failing_compilation_first");
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.clone(), rhs.clone(), out.clone()];
let uid = fresh_tune_key_uid();
let test_set = TUNER.init(&"test".to_string(), move || {
let client = test_client(&DummyDevice);
let shapes = vec![vec![1, 3], vec![1, 3], vec![1, 3]];
dummy::addition_set_with_failing_compilation_first(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(lhs).unwrap().to_vec(), vec![0, 1, 2]);
assert_eq!(client.read_one(rhs).unwrap().to_vec(), vec![4, 4, 4]);
assert_eq!(client.read_one(out).unwrap().to_vec(), vec![4, 5, 6]);
}
#[test_log::test]
#[cfg(all(feature = "std", not(target_family = "wasm")))]
#[serial_test::serial]
fn autotune_stops_sampling_an_eliminated_candidate() {
use cubecl_server::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(&"test".to_string(), 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_server::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_server::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_server::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(&"test".to_string(), || {
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"
);
}
#[test_log::test]
#[cfg(feature = "std")]
#[serial_test::serial]
fn a_dry_run_reserves_without_mapping() {
use cubecl_server::dry_run::DryRun;
use cubecl_server::memory_management::MemoryPoolReport;
let client = test_client(&DummyDevice);
let dry_run = DryRun::new();
const SIZE: u64 = 32 * 1024 * 1024;
fn arena(report: &cubecl_server::memory_management::MemoryReport) -> MemoryPoolReport {
report
.dynamic
.iter()
.find(|pool| pool.largest_alloc == SIZE)
.expect("some pool served the buffer")
.clone()
}
let out = client.empty(SIZE as usize);
client.launch(
Box::new(KernelTask::new(DummyElementwiseAddition)),
CubeCount::Static(1, 1, 1),
KernelArguments::new().with_buffers(vec![
out.clone().binding(),
out.clone().binding(),
out.clone().binding(),
]),
);
let report = client.memory_report();
let pool = arena(&report);
assert_eq!(
pool.pages_unmapped, pool.pages,
"the reservation must have no device backing: {report:?}"
);
assert!(pool.pages >= 1, "{report:?}");
let data = client.read_one(out).unwrap();
assert_eq!(data.len(), SIZE as usize);
let report = client.memory_report();
assert_eq!(
arena(&report).pages_unmapped,
0,
"resolution installed the backing: {report:?}"
);
drop(dry_run);
}
#[test_log::test]
#[cfg(feature = "std")]
#[serial_test::serial]
fn a_set_is_built_once_per_device_not_once_per_process() {
use std::sync::atomic::{AtomicUsize, Ordering};
static TUNER: LocalTuner<String, String> =
local_tuner!("a_set_is_built_once_per_device_not_once_per_process");
static BUILDS: AtomicUsize = AtomicUsize::new(0);
let client = test_client(&DummyDevice);
let build = |device: &str| {
TUNER.init(&device.to_string(), || {
BUILDS.fetch_add(1, Ordering::Relaxed);
let shapes = vec![vec![1, 3], vec![1, 3], vec![1, 3]];
dummy::addition_set(test_client(&DummyDevice), shapes)
})
};
let first = build("gpu-0");
assert_eq!(BUILDS.load(Ordering::Relaxed), 1);
let first_again = build("gpu-0");
assert_eq!(
BUILDS.load(Ordering::Relaxed),
1,
"the same device must reuse its set rather than rebuild it"
);
assert!(
std::sync::Arc::ptr_eq(&first, &first_again),
"the same device must get the very same set back"
);
let second = build("gpu-1");
assert_eq!(
BUILDS.load(Ordering::Relaxed),
2,
"a device that has not tuned yet must build its own set"
);
assert!(
!std::sync::Arc::ptr_eq(&first, &second),
"one device's set must not answer for another's"
);
let lhs = client.create_from_slice(&[0, 1, 2]);
let rhs = client.create_from_slice(&[4, 4, 4]);
let out = client.empty(3);
TUNER.execute(
&"gpu-1".to_string(),
&client,
second,
vec![lhs, rhs, out.clone()],
);
assert_eq!(client.read_one(out).unwrap().to_vec(), Vec::from([4, 5, 6]));
assert_eq!(BUILDS.load(Ordering::Relaxed), 2);
}
#[test_log::test]
#[cfg(all(feature = "std", persistence))]
#[serial_test::serial]
fn a_compilation_is_recorded_with_its_outcome() {
use cubecl_environment::persistence::Database;
use cubecl_environment::records::{RecordLevel, Records};
use cubecl_server::compiler::{CompilationOutcome, CompilationRecord, CompilationRecording};
use cubecl_server::id::KernelId;
struct Recorded;
let root = tempfile::tempdir().unwrap();
rooted_at(root.path());
recording_at(RecordLevel::Basic);
let id = KernelId::new::<Recorded>().info(3u32);
let stored = true;
CompilationRecording::new(&id).compiled(stored);
CompilationRecording::new(&id).loaded();
CompilationRecording::new(&id).rekeyed(stored);
let database = Database::open_active().unwrap();
let trips = Records::new(&database).read::<CompilationRecord>();
let outcomes: Vec<CompilationOutcome> = trips.iter().map(|trip| trip.record.outcome).collect();
assert_eq!(
outcomes,
vec![
CompilationOutcome::Compiled,
CompilationOutcome::Loaded,
CompilationOutcome::Rekeyed
]
);
assert!(trips[0].record.kernel.ends_with("Recorded"));
assert_eq!(trips[0].record.key, trips[1].record.key);
}
#[test_log::test]
#[cfg(all(feature = "std", persistence))]
#[serial_test::serial]
fn a_compilation_keeps_its_code_only_when_records_are_full() {
use cubecl_environment::persistence::Database;
use cubecl_environment::records::{RecordLevel, Records};
use cubecl_server::compiler::{CompilationRecord, CompilationRecording};
use cubecl_server::id::KernelId;
struct Coded;
let root = tempfile::tempdir().unwrap();
rooted_at(root.path());
let id = KernelId::new::<Coded>();
let stored = true;
for level in [RecordLevel::Basic, RecordLevel::Full] {
recording_at(level);
let mut recording = CompilationRecording::new(&id);
recording.source("source");
recording.compiled(stored);
}
recording_at(RecordLevel::Basic);
let database = Database::open_active().unwrap();
let sources: Vec<Option<String>> = Records::new(&database)
.read::<CompilationRecord>()
.into_iter()
.map(|trip| trip.record.source)
.collect();
assert_eq!(sources, vec![None, Some("source".to_string())]);
}
#[test_log::test]
#[cfg(all(feature = "std", persistence))]
#[serial_test::serial]
fn a_compile_nothing_stored_leaves_no_session() {
use cubecl_environment::persistence::{Database, Namespace, Store, StoreOptions};
use cubecl_environment::records::{RecordLevel, Records};
use cubecl_server::compiler::{CompilationRecord, CompilationRecording, store_compiled};
use cubecl_server::id::KernelId;
struct Unstored;
let root = tempfile::tempdir().unwrap();
rooted_at(root.path());
recording_at(RecordLevel::Basic);
let mut store: Store<u32, u32> =
Store::new(StoreOptions::new().storage(Namespace::new("test/compiled")));
assert!(store_compiled(&mut store, 1, 1), "the store took it");
let id = KernelId::new::<Unstored>();
let stored = false;
CompilationRecording::new(&id).compiled(stored);
let database = Database::open_active().unwrap();
let records = Records::new(&database);
assert!(records.sessions().is_empty());
assert!(records.read::<CompilationRecord>().is_empty());
}
#[test_log::test]
#[cfg(all(feature = "std", persistence))]
#[serial_test::serial]
fn a_memory_snapshot_is_recorded_under_its_label() {
use cubecl_environment::persistence::Database;
use cubecl_environment::records::{RecordLevel, Records};
use cubecl_server::memory_management::MemoryRecord;
let root = tempfile::tempdir().unwrap();
rooted_at(root.path());
recording_at(RecordLevel::Basic);
let client = test_client(&DummyDevice);
let _held = client.create_from_slice(&[1, 2, 3]);
client.record_memory("model loaded");
let database = Database::open_active().unwrap();
assert!(
Records::new(&database).read::<MemoryRecord>().is_empty(),
"held until the session changes something"
);
struct Stored;
let id = cubecl_server::id::KernelId::new::<Stored>();
let stored = true;
cubecl_server::compiler::CompilationRecording::new(&id).compiled(stored);
let snapshots = Records::new(&database).read::<MemoryRecord>();
assert_eq!(snapshots.len(), 1);
assert_eq!(snapshots[0].record.label, "model loaded");
assert_eq!(snapshots[0].record.report, client.memory_report());
}
#[test_log::test]
#[serial_test::serial]
fn a_launch_is_collected_while_a_collection_is_open() {
use cubecl_server::launched::LaunchedKernels;
let client = test_client(&DummyDevice);
let launch = || {
let lhs = client.create_from_slice(&[0, 1, 2]);
let rhs = client.create_from_slice(&[4, 4, 4]);
let out = client.empty(3);
let kernel = KernelTask::new(DummyElementwiseAddition);
let id = cubecl_server::kernel::KernelMetadata::id(&kernel);
client.launch(
Box::new(kernel),
CubeCount::Static(1, 1, 1),
KernelArguments::new().with_buffers(vec![lhs.binding(), rhs.binding(), out.binding()]),
);
id
};
let collection = LaunchedKernels::new();
let id = launch();
let launched = collection.finish();
assert!(launched.contains(&id.stable_hash()));
let collection = LaunchedKernels::new();
let launched = collection.finish();
assert!(launched.is_empty(), "a new collection starts empty");
}