use cubecl_environment::sync::{AtomicUsize, Ordering};
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum LaunchMode {
Execute,
Skip,
}
impl LaunchMode {
pub fn is_skipped(self) -> bool {
matches!(self, LaunchMode::Skip)
}
}
pub fn launch_mode() -> LaunchMode {
if !dry_run() || real_run::depth() > 0 {
return LaunchMode::Execute;
}
LaunchMode::Skip
}
static DRY_RUN: AtomicUsize = AtomicUsize::new(0);
pub fn dry_run() -> bool {
DRY_RUN.load(Ordering::Relaxed) > 0
}
#[derive(Debug)]
pub struct DryRun {
_private: (),
}
impl DryRun {
#[allow(clippy::new_without_default, reason = "a guard is not a value")]
pub fn new() -> Self {
DRY_RUN.fetch_add(1, Ordering::Relaxed);
Self { _private: () }
}
}
impl Drop for DryRun {
fn drop(&mut self) {
DRY_RUN.fetch_sub(1, Ordering::Relaxed);
}
}
#[derive(Debug)]
pub struct RealRun {
_private: (),
}
impl RealRun {
#[allow(clippy::new_without_default, reason = "a guard is not a value")]
pub fn new() -> Self {
real_run::enter();
Self { _private: () }
}
}
impl Drop for RealRun {
fn drop(&mut self) {
real_run::exit();
}
}
#[cfg(feature = "std")]
mod real_run {
use core::cell::Cell;
std::thread_local! {
static DEPTH: Cell<usize> = const { Cell::new(0) };
}
pub(super) fn depth() -> usize {
DEPTH.with(|depth| depth.get())
}
pub(super) fn enter() {
DEPTH.with(|depth| depth.set(depth.get() + 1));
}
pub(super) fn exit() {
DEPTH.with(|depth| depth.set(depth.get().saturating_sub(1)));
}
}
#[cfg(not(feature = "std"))]
mod real_run {
pub(super) fn depth() -> usize {
0
}
pub(super) fn enter() {}
pub(super) fn exit() {}
}
#[cfg(test)]
mod tests {
use super::*;
use alloc::vec;
#[test]
fn real_run_nests() {
assert_eq!(real_run::depth(), 0);
let outer = RealRun::new();
{
let _inner = RealRun::new();
assert_eq!(real_run::depth(), 2);
}
assert_eq!(real_run::depth(), 1, "the outer guard is still open");
drop(outer);
assert_eq!(real_run::depth(), 0);
}
#[test]
#[serial_test::serial]
fn launches_execute_by_default() {
assert_eq!(launch_mode(), LaunchMode::Execute);
let _real_run = RealRun::new();
assert_eq!(launch_mode(), LaunchMode::Execute);
}
#[test]
#[serial_test::serial]
fn a_dry_run_spares_the_measurements() {
let _dry_run = DryRun::new();
assert_eq!(launch_mode(), LaunchMode::Skip);
{
let _real_run = RealRun::new();
assert_eq!(launch_mode(), LaunchMode::Execute, "a measurement runs");
}
assert_eq!(launch_mode(), LaunchMode::Skip);
}
#[test]
#[serial_test::serial]
fn dry_runs_nest() {
assert!(!dry_run());
{
let _outer = DryRun::new();
{
let _inner = DryRun::new();
assert!(dry_run());
}
assert!(dry_run(), "the outer guard is still in force");
}
assert!(!dry_run(), "and the process is back to executing");
}
}