use crate::{CubeDevice, tensor::CubeTensor};
use burn_backend::cubecl::dtype_to_storage_type;
use burn_backend::{
Backend, BackendGraph, BackendTypes, DTypeUsage, DTypeUsageSet, ExecutionError,
InstallMemoryPoolsError, MemoryPoolLayout, MemoryPoolUsage, ProfileDuration, ProfileOptions,
ProfileToken, SlicedPool, SlicedPoolReport, TensorData, profile_with_tokens,
};
use burn_std::{BoolStore, DType, id::StreamId, quantization::quantizable};
use cubecl::device::DeviceId;
use cubecl::{
MemoryConfiguration, MemoryPoolKind,
client::{Client, ProfileWindow},
config::memory::{MemoryPoolConfig, MemoryPoolsConfig, MemoryPoolsPreset},
config::size::MemorySize,
features::{MmaConfig, TypeUsage},
ir::ElemType,
server::{ProfileError, ProfilingToken},
};
#[cfg(not(feature = "fusion"))]
use burn_backend::tensor::{BoolTensor, FloatTensor, IntTensor, QuantizedTensor};
#[cfg(not(feature = "fusion"))]
use burn_ir::{BackendIr, TensorHandle};
fn qfloat_params_usable(client: &Client, dtype: DType) -> bool {
let DType::QFloat(scheme) = dtype else {
return true;
};
quantizable(&scheme)
&& client
.properties()
.type_usage(ElemType::from_scale_dtype(scheme.scale_dtype()))
.is_superset(TypeUsage::Buffer | TypeUsage::Conversion)
}
fn graph_err(err: impl core::fmt::Display) -> ExecutionError {
ExecutionError::WithContext {
reason: format!("{err}"),
}
}
fn profile_err(err: ProfileError) -> ExecutionError {
ExecutionError::WithContext {
reason: format!("{err}"),
}
}
fn empty_window() -> ProfileDuration {
ProfileDuration::new_device_time_maybe(async move { None })
}
#[derive(Clone, Debug)]
pub struct CubeGraph {
graph: cubecl::client::Graph,
device: CubeDevice,
}
#[derive(new)]
pub struct CubeBackend;
impl BackendTypes for CubeBackend {
type Device = CubeDevice;
type FloatTensorPrimitive = CubeTensor;
type IntTensorPrimitive = CubeTensor;
type BoolTensorPrimitive = CubeTensor;
type QuantizedTensorPrimitive = CubeTensor;
type GraphPrimitive = CubeGraph;
}
impl Backend for CubeBackend {
fn name(device: &Self::Device) -> String {
let client = device.client();
format!("cubecl<{}>", client.name())
}
fn seed(_device: &Self::Device, seed: u64) {
cubek::random::seed(seed);
}
fn ad_enabled(_device: &Self::Device) -> bool {
false
}
fn sync(device: &Self::Device) -> Result<(), ExecutionError> {
let client = device.client();
futures_lite::future::block_on(client.sync()).map_err(|err| ExecutionError::WithContext {
reason: format!("{err}"),
})
}
fn profile<O: Send + 'static>(
device: &Self::Device,
options: ProfileOptions,
func: impl FnOnce() -> O + Send,
) -> Result<(O, ProfileDuration), ExecutionError> {
profile_with_tokens::<Self, O>(device, options, func)
}
fn profile_start(device: &Self::Device) -> Result<Option<ProfileToken>, ExecutionError> {
let client = device.client();
client
.profile_start()
.map(|window| {
Some(ProfileToken {
id: window.token.id,
opened_on: window.stream_id.value,
})
})
.map_err(profile_err)
}
fn profile_end(
device: &Self::Device,
token: ProfileToken,
_options: ProfileOptions,
) -> Result<ProfileDuration, ExecutionError> {
let client = device.client();
let window = ProfileWindow {
stream_id: StreamId {
value: token.opened_on,
},
token: ProfilingToken { id: token.id },
};
match client.profile_end(window) {
Ok(duration) => Ok(duration),
Err(ProfileError::NotMeasured { .. }) => Ok(empty_window()),
Err(err) => Err(profile_err(err)),
}
}
fn profile_abandon(device: &Self::Device, token: ProfileToken) {
let window = ProfileWindow {
stream_id: StreamId {
value: token.opened_on,
},
token: ProfilingToken { id: token.id },
};
device.client().profile_abandon(window);
}
fn graph_prepare(device: &Self::Device) -> Result<(), ExecutionError> {
let client = device.client();
client.graph_prepare().map_err(graph_err)
}
fn graph_start_capture(device: &Self::Device) -> Result<(), ExecutionError> {
let client = device.client();
client.start_capture().map_err(graph_err)
}
fn graph_stop_capture(device: &Self::Device) -> Result<BackendGraph<Self>, ExecutionError> {
let client = device.client();
let graph = client.stop_capture().map_err(graph_err)?;
Ok(CubeGraph {
graph,
device: device.clone(),
})
}
unsafe fn graph_replay(
device: &Self::Device,
graph: &BackendGraph<Self>,
) -> Result<(), ExecutionError> {
if &graph.device != device {
return Err(ExecutionError::WithContext {
reason: format!(
"The graph was captured on {:?} and cannot replay on {device:?}",
graph.device
),
});
}
unsafe { graph.graph.replay() }.map_err(graph_err)
}
fn memory_persistent_allocations<
Output: Send,
Input: Send,
Func: Fn(Input) -> Output + Send,
>(
device: &Self::Device,
input: Input,
func: Func,
) -> Output {
let client = device.client();
client.memory_persistent_allocation(input, func)
}
fn memory_cleanup(device: &Self::Device) {
let client = device.client();
client.memory_cleanup();
}
fn memory_install_pools(
device: &Self::Device,
layout: MemoryPoolLayout,
) -> Result<(), InstallMemoryPoolsError> {
let client = device.client();
let properties = &client.properties().memory;
let config = pool_config(layout, properties.alignment.max(1))?;
MemoryConfiguration::default()
.resolve(Some(&config), properties)
.map_err(|err| InstallMemoryPoolsError::InvalidLayout {
reason: err.to_string(),
})?;
client
.install_memory_pools(&config)
.map_err(runtime_install_error)
}
fn memory_pool_report(device: &Self::Device) -> Option<Vec<SlicedPoolReport>> {
let report = device.client().memory_report();
Some(
report
.dynamic
.iter()
.filter_map(|pool| match pool.kind {
MemoryPoolKind::Sliced { page_size, .. } => Some(SlicedPoolReport {
page_size,
pages: pool.pages,
pages_peak: pool.pages_peak,
largest_alloc: pool.largest_alloc,
}),
MemoryPoolKind::Direct => Some(SlicedPoolReport {
page_size: 0,
pages: pool.pages,
pages_peak: pool.pages_peak,
largest_alloc: pool.largest_alloc,
}),
_ => None,
})
.collect(),
)
}
fn memory_pool_usage(device: &Self::Device) -> Option<MemoryPoolUsage> {
let usage = device.client().memory_usage();
Some(MemoryPoolUsage {
number_allocs: usage.number_allocs,
bytes_in_use: usage.bytes_in_use,
bytes_padding: usage.bytes_padding,
bytes_reserved: usage.bytes_reserved,
})
}
fn staging<'a, Iter>(data: Iter, device: &Self::Device)
where
Iter: Iterator<Item = &'a mut TensorData>,
{
let client = device.client();
client.staging(data.map(|td| &mut td.bytes), false);
}
fn supports_dtype(device: &Self::Device, dtype: DType) -> bool {
if let DType::Bool(BoolStore::Native) = dtype {
return false;
}
let client = device.client();
if !qfloat_params_usable(&client, dtype) {
return false;
}
let type_usage = client.properties().type_usage(dtype_to_storage_type(dtype));
type_usage.is_superset(
TypeUsage::Buffer
| TypeUsage::Conversion
| TypeUsage::Arithmetic
| TypeUsage::DotProduct,
)
}
fn dtype_usage(device: &Self::Device, dtype: DType) -> DTypeUsageSet {
if let DType::Bool(BoolStore::Native) = dtype {
return DTypeUsageSet::empty();
}
let client = device.client();
if !qfloat_params_usable(&client, dtype) {
return DTypeUsageSet::empty();
}
let props = client.properties();
let storage = dtype_to_storage_type(dtype);
let usage = props.type_usage(storage);
let mut out = DTypeUsageSet::new();
if usage.is_superset(TypeUsage::Buffer | TypeUsage::Conversion) {
out |= DTypeUsage::Storage;
}
if usage.contains(TypeUsage::Arithmetic) {
out |= DTypeUsage::Arithmetic;
}
let has_mma = |cfg: &MmaConfig| {
cfg.a_type == storage || cfg.b_type == storage || cfg.cd_type == storage
};
if props.features.matmul.cmma.iter().any(has_mma)
|| props.features.matmul.mma.iter().any(has_mma)
{
out |= DTypeUsage::Accelerated;
}
out
}
fn device_count(type_id: u16) -> usize {
CubeDevice::enumerate(DeviceId::new(type_id, 0)).len()
}
fn flush(device: &Self::Device) {
let client = device.client();
client.flush().unwrap();
}
}
impl core::fmt::Debug for CubeBackend {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.write_str("CubeCLBackend")
}
}
impl Clone for CubeBackend {
fn clone(&self) -> Self {
Self::new()
}
}
impl Default for CubeBackend {
fn default() -> Self {
Self::new()
}
}
#[cfg(not(feature = "fusion"))]
impl BackendIr for CubeBackend {
type Handle = CubeTensor;
fn float_tensor(handle: TensorHandle<Self::Handle>) -> FloatTensor<Self> {
handle.handle
}
fn int_tensor(handle: TensorHandle<Self::Handle>) -> IntTensor<Self> {
handle.handle
}
fn bool_tensor(handle: TensorHandle<Self::Handle>) -> BoolTensor<Self> {
handle.handle
}
fn quantized_tensor(handle: TensorHandle<Self::Handle>) -> QuantizedTensor<Self> {
handle.handle
}
fn float_tensor_handle(tensor: FloatTensor<Self>) -> Self::Handle {
tensor
}
fn int_tensor_handle(tensor: IntTensor<Self>) -> Self::Handle {
tensor
}
fn bool_tensor_handle(tensor: BoolTensor<Self>) -> Self::Handle {
tensor
}
fn quantized_tensor_handle(tensor: QuantizedTensor<Self>) -> Self::Handle {
tensor
}
}
fn pool_config(
layout: MemoryPoolLayout,
alignment: u64,
) -> Result<MemoryPoolsConfig, InstallMemoryPoolsError> {
let config = match layout {
MemoryPoolLayout::Sliced(pools) => MemoryPoolsConfig::Explicit(
pools
.into_iter()
.map(
|SlicedPool {
page_size,
pages,
max_slice,
}| {
let page_size = align_up(page_size, alignment)?;
let max_pool_size = pages
.map(|pages| {
page_size.checked_mul(pages).ok_or_else(|| {
InstallMemoryPoolsError::InvalidLayout {
reason: format!(
"a cap of {pages} pages of {page_size} B overflows"
),
}
})
})
.transpose()?;
Ok(MemoryPoolConfig::Sliced {
page_size: MemorySize(page_size),
max_slice_size: max_slice
.map(|size| align_up(size, alignment))
.transpose()?
.map(MemorySize),
max_pool_size: max_pool_size.map(MemorySize),
dealloc_period: None,
})
},
)
.collect::<Result<Vec<_>, _>>()?,
),
MemoryPoolLayout::Direct => {
MemoryPoolsConfig::Explicit(vec![MemoryPoolConfig::Direct { reclaim_at: None }])
}
MemoryPoolLayout::SubSlices => MemoryPoolsConfig::Preset(MemoryPoolsPreset::SubSlices),
MemoryPoolLayout::ExclusivePages => {
MemoryPoolsConfig::Preset(MemoryPoolsPreset::ExclusivePages)
}
};
Ok(config)
}
fn align_up(size: u64, alignment: u64) -> Result<u64, InstallMemoryPoolsError> {
size.checked_next_multiple_of(alignment)
.ok_or_else(|| InstallMemoryPoolsError::InvalidLayout {
reason: format!("{size} B cannot be aligned up to {alignment} B"),
})
}
fn runtime_install_error(err: cubecl::InstallMemoryPoolsError) -> InstallMemoryPoolsError {
match err {
cubecl::InstallMemoryPoolsError::PoolsInUse { bytes_in_use } => {
InstallMemoryPoolsError::PoolsInUse { bytes_in_use }
}
cubecl::InstallMemoryPoolsError::Unsupported => InstallMemoryPoolsError::Unsupported,
}
}