use std::fmt::{Display, Write};
use cubecl_core::ir::BarrierLevel;
use crate::shared::FmtLeft;
use super::{Component, Dialect, Value};
#[derive(Debug, Clone)]
pub enum BarrierOps<D: Dialect> {
Init {
barrier: Value<D>,
is_elected: Value<D>,
arrival_count: Value<D>,
level: BarrierLevel,
},
InitManual {
barrier: Value<D>,
arrival_count: Value<D>,
},
MemCopyAsync {
barrier: Value<D>,
source: Value<D>,
destination: Value<D>,
source_length: Value<D>,
cooperative: bool,
},
MemCopyAsyncTx {
barrier: Value<D>,
source: Value<D>,
destination: Value<D>,
source_length: Value<D>,
},
CopyAsync {
source: Value<D>,
destination: Value<D>,
source_length: Value<D>,
copy_size: u32,
checked: bool,
},
MemCopyAsyncTensorGlobalToShared {
barrier: Value<D>,
smem_buffer: Value<D>,
tensor_map: Value<D>,
indices: Vec<Value<D>>,
},
TmaLoadIm2col {
barrier: Value<D>,
smem_buffer: Value<D>,
tensor_map: Value<D>,
indices: Vec<Value<D>>,
offsets: Vec<Value<D>>,
},
Arrive {
barrier: Value<D>,
token: Value<D>,
},
ArriveTx {
barrier: Value<D>,
token: Value<D>,
arrive_count_update: Value<D>,
transaction_count_update: Value<D>,
},
ArriveCopyAsync {
barrier: Value<D>,
},
ExpectTx {
barrier: Value<D>,
transaction_count_update: Value<D>,
},
Wait {
barrier: Value<D>,
token: Value<D>,
},
WaitParity {
barrier: Value<D>,
phase: Value<D>,
},
ArriveAndWait {
barrier: Value<D>,
level: BarrierLevel,
},
}
impl<D: Dialect> BarrierOps<D> {
pub fn barrier_id(&self) -> u32 {
match self {
BarrierOps::MemCopyAsync { barrier, .. } => barrier.id().unwrap(),
BarrierOps::MemCopyAsyncTx { barrier, .. } => barrier.id().unwrap(),
BarrierOps::Init { barrier, .. } => barrier.id().unwrap(),
BarrierOps::InitManual { barrier, .. } => barrier.id().unwrap(),
BarrierOps::ArriveAndWait { barrier, .. } => barrier.id().unwrap(),
BarrierOps::Arrive { barrier, .. } => barrier.id().unwrap(),
BarrierOps::ArriveTx { barrier, .. } => barrier.id().unwrap(),
BarrierOps::Wait { barrier, .. } => barrier.id().unwrap(),
BarrierOps::WaitParity { barrier, .. } => barrier.id().unwrap(),
BarrierOps::MemCopyAsyncTensorGlobalToShared { barrier, .. } => barrier.id().unwrap(),
BarrierOps::TmaLoadIm2col { barrier, .. } => barrier.id().unwrap(),
BarrierOps::ExpectTx { barrier, .. } => barrier.id().unwrap(),
BarrierOps::CopyAsync { .. } => 0,
BarrierOps::ArriveCopyAsync { barrier } => barrier.id().unwrap(),
}
}
}
impl<D: Dialect> Display for BarrierOps<D> {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
match self {
BarrierOps::Init {
barrier,
is_elected,
arrival_count,
level,
} => match level {
BarrierLevel::Unit => write!(
f,
"
init(&{barrier}, {arrival_count});
"
),
BarrierLevel::Cube => write!(
f,
"
if ({is_elected}) {{
init(&{barrier}, {arrival_count});
}}
__syncthreads();
",
),
},
BarrierOps::InitManual {
barrier,
arrival_count,
} => {
writeln!(f, "init(&{barrier}, {arrival_count});")
}
BarrierOps::MemCopyAsync {
barrier,
source,
destination,
source_length,
cooperative,
} => {
let item = source.item();
let size = format!("sizeof({})", item.value_ty());
match cooperative {
false => write!(
f,
"
cuda::memcpy_async({destination}, {source}, {source_length} * {size}, {barrier});
"
),
true => write!(
f,
"
cuda::memcpy_async(thread_block, {destination}, {source}, {source_length} * {size}, {barrier});
"
),
}
}
BarrierOps::MemCopyAsyncTx {
barrier,
source,
destination,
source_length,
} => {
let item = source.item();
let size = format!("sizeof({})", item.value_ty());
write!(
f,
"
cuda::device::memcpy_async_tx({destination}, {source}, {source_length} * {size}, {barrier});
"
)
}
BarrierOps::CopyAsync {
source,
destination,
source_length,
copy_size,
checked,
} => {
let item = source.item();
let size = format!("{source_length} * sizeof({})", item.value_ty());
match *checked {
false => write!(
f,
"
__cp_async_shared_global<{copy_size}>({source}, {destination});
"
),
true => write!(
f,
"
__cp_async_shared_global<{copy_size}>({source}, {destination}, {size});
"
),
}
}
BarrierOps::MemCopyAsyncTensorGlobalToShared {
barrier,
smem_buffer,
tensor_map,
indices,
} => {
let rank = indices.len();
let smem_ptr = smem_buffer.fmt_ptr();
let indices = indices.iter().rev().fold(String::new(), |mut s, it| {
let _ = write!(s, "{it}, ");
s
});
writeln!(
f,
"cuda::device::experimental::cp_async_bulk_tensor_{rank}d_global_to_shared({smem_ptr}, &{tensor_map}, {indices} {barrier});"
)
}
BarrierOps::TmaLoadIm2col {
barrier,
smem_buffer,
tensor_map,
indices,
offsets,
} => {
let rank = indices.len();
let smem_ptr = smem_buffer.fmt_ptr();
let args: Vec<_> = indices
.iter()
.rev()
.map(|it| it.to_string())
.chain(offsets.iter().rev().map(|it| it.to_string()))
.collect();
writeln!(
f,
"tma_load_im2col_{rank}d(&{tensor_map}, {barrier}, {smem_ptr}, {});",
args.join(", ")
)
}
BarrierOps::Arrive { barrier, token, .. } => {
writeln!(f, "{} = {barrier}.arrive();", token.fmt_left())
}
BarrierOps::ArriveTx {
barrier,
token,
arrive_count_update,
transaction_count_update,
} => {
writeln!(
f,
"{} = cuda::device::barrier_arrive_tx({barrier}, {arrive_count_update}, {transaction_count_update});",
token.fmt_left()
)
}
BarrierOps::ArriveCopyAsync { barrier } => {
writeln!(f, "__cp_async_arrive({barrier});")
}
BarrierOps::ExpectTx {
barrier,
transaction_count_update,
} => {
writeln!(
f,
"cuda::device::barrier_expect_tx({barrier}, {transaction_count_update});"
)
}
BarrierOps::Wait { barrier, token } => {
writeln!(f, "{barrier}.wait(std::move({token}));")
}
BarrierOps::WaitParity { barrier, phase } => {
writeln!(f, "{barrier}.wait_parity({phase});")
}
BarrierOps::ArriveAndWait { barrier, .. } => {
writeln!(f, "{barrier}.arrive_and_wait();")
}
}
}
}