cubecl-cpp 0.11.0-pre.1

CPP transpiler for CubeCL
Documentation
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();")
            }
        }
    }
}