Skip to main content

burn_cubecl/ops/
bool_tensor.rs

1use crate::{
2    CubeBackend, CubeRuntime,
3    element::BoolElement,
4    kernel::{self, AndOp, OrOp},
5    tensor::CubeTensor,
6};
7use burn_backend::cubecl::dtype_to_storage_type;
8use burn_backend::{
9    ExecutionError, Slice,
10    ops::BoolTensorOps,
11    tensor::{BoolTensor, Device, FloatTensor, IntTensor},
12};
13use burn_backend::{Scalar, Shape, TensorData};
14use burn_std::{BoolDType, BoolStore, DType, FloatDType, IntDType};
15use cubecl::prelude::InputScalar;
16use cubek::reduce::components::instructions::ReduceOperationConfig;
17use std::ops::Range;
18
19use super::{expand, numeric, permute, unfold};
20
21/// The boolean storage of a cubecl bool tensor. Cubecl backends never use
22/// native bool (see `CubeBackend::supports_dtype`), so it is always `U8`/`U32`.
23fn bool_store<R: CubeRuntime>(tensor: &CubeTensor<R>) -> BoolDType {
24    match tensor.dtype {
25        DType::Bool(store) => store,
26        other => unreachable!("cubecl bool tensors are always Bool(_): {other:?}"),
27    }
28}
29
30impl<R: CubeRuntime> BoolTensorOps<Self> for CubeBackend<R> {
31    fn bool_empty(shape: Shape, device: &Device<Self>, dtype: BoolDType) -> BoolTensor<Self> {
32        super::empty(shape, device, dtype.into())
33    }
34
35    fn bool_zeros(shape: Shape, device: &Device<Self>, dtype: BoolDType) -> BoolTensor<Self> {
36        numeric::zeros(device.clone(), shape, dtype.into())
37    }
38
39    fn bool_ones(shape: Shape, device: &Device<Self>, dtype: BoolDType) -> BoolTensor<Self> {
40        numeric::ones(device.clone(), shape, dtype.into())
41    }
42
43    async fn bool_into_data(tensor: BoolTensor<Self>) -> Result<TensorData, ExecutionError> {
44        super::into_data(tensor).await
45    }
46
47    fn bool_from_data(data: TensorData, device: &Device<Self>) -> BoolTensor<Self> {
48        if !matches!(
49            data.dtype,
50            DType::Bool(BoolStore::U8) | DType::Bool(BoolStore::U32)
51        ) {
52            unimplemented!("Unsupported dtype for `bool_from_data` {:?}", data.dtype);
53        }
54        super::from_data(data, device)
55    }
56
57    fn bool_into_int(tensor: BoolTensor<Self>, out_dtype: IntDType) -> IntTensor<Self> {
58        kernel::bool_cast(tensor, out_dtype.into())
59    }
60
61    fn bool_to_device(tensor: BoolTensor<Self>, device: &Device<Self>) -> BoolTensor<Self> {
62        super::to_device(tensor, device)
63    }
64
65    fn bool_reshape(tensor: BoolTensor<Self>, shape: Shape) -> BoolTensor<Self> {
66        super::reshape(tensor, shape)
67    }
68
69    fn bool_slice(tensor: BoolTensor<Self>, slices: &[Slice]) -> BoolTensor<Self> {
70        // Check if all steps are 1
71        let all_steps_one = slices.iter().all(|info| info.step == 1);
72
73        if all_steps_one {
74            // Use optimized slice for step=1
75            let simple_ranges: Vec<Range<usize>> = slices
76                .iter()
77                .enumerate()
78                .map(|(i, slice)| slice.to_range(tensor.meta.shape()[i]))
79                .collect();
80
81            kernel::slice(tensor, &simple_ranges)
82        } else {
83            // Use slice with steps kernel
84            kernel::slice_with_steps(tensor, slices)
85        }
86    }
87
88    fn bool_slice_assign(
89        tensor: BoolTensor<Self>,
90        ranges: &[Slice],
91        value: BoolTensor<Self>,
92    ) -> BoolTensor<Self> {
93        kernel::slice_assign(tensor, ranges, value)
94    }
95
96    fn bool_equal(lhs: BoolTensor<Self>, rhs: BoolTensor<Self>) -> BoolTensor<Self> {
97        let dtype = lhs.dtype;
98        kernel::equal(lhs, rhs, dtype)
99    }
100
101    fn bool_not(tensor: BoolTensor<Self>) -> BoolTensor<Self> {
102        let dtype = tensor.dtype;
103        let storage = dtype_to_storage_type(dtype);
104        let scalar = match dtype {
105            DType::Bool(BoolStore::U32) => InputScalar::new(u32::false_val(), storage),
106            DType::Bool(BoolStore::U8) => InputScalar::new(u8::false_val(), storage),
107            other => unimplemented!("Unsupported dtype for `bool_from_data` {other:?}"),
108        };
109        kernel::equal_elem(tensor, scalar, dtype)
110    }
111
112    fn bool_and(lhs: BoolTensor<Self>, rhs: BoolTensor<Self>) -> BoolTensor<Self> {
113        kernel::launch_binop::<R, AndOp>(lhs, rhs)
114    }
115
116    fn bool_or(lhs: BoolTensor<Self>, rhs: BoolTensor<Self>) -> BoolTensor<Self> {
117        kernel::launch_binop::<R, OrOp>(lhs, rhs)
118    }
119
120    fn bool_any(tensor: BoolTensor<Self>) -> BoolTensor<Self> {
121        let store = bool_store(&tensor);
122        kernel::reduce::reduce_logical(tensor, None, ReduceOperationConfig::Any, store)
123    }
124
125    fn bool_any_dim(tensor: BoolTensor<Self>, dim: usize) -> BoolTensor<Self> {
126        let store = bool_store(&tensor);
127        kernel::reduce::reduce_logical(tensor, Some(dim), ReduceOperationConfig::Any, store)
128    }
129
130    fn bool_all(tensor: BoolTensor<Self>) -> BoolTensor<Self> {
131        let store = bool_store(&tensor);
132        kernel::reduce::reduce_logical(tensor, None, ReduceOperationConfig::All, store)
133    }
134
135    fn bool_all_dim(tensor: BoolTensor<Self>, dim: usize) -> BoolTensor<Self> {
136        let store = bool_store(&tensor);
137        kernel::reduce::reduce_logical(tensor, Some(dim), ReduceOperationConfig::All, store)
138    }
139
140    fn bool_into_float(tensor: BoolTensor<Self>, out_dtype: FloatDType) -> FloatTensor<Self> {
141        kernel::bool_cast(tensor, out_dtype.into())
142    }
143
144    fn bool_swap_dims(mut tensor: BoolTensor<Self>, dim1: usize, dim2: usize) -> BoolTensor<Self> {
145        tensor.meta.swap(dim1, dim2);
146
147        tensor
148    }
149
150    fn bool_repeat_dim(tensor: BoolTensor<Self>, dim: usize, times: usize) -> BoolTensor<Self> {
151        kernel::repeat_dim(tensor, dim, times)
152    }
153
154    fn bool_permute(tensor: BoolTensor<Self>, axes: &[usize]) -> BoolTensor<Self> {
155        permute(tensor, axes)
156    }
157
158    fn bool_expand(tensor: BoolTensor<Self>, shape: Shape) -> BoolTensor<Self> {
159        expand(tensor, shape)
160    }
161
162    fn bool_select(
163        tensor: BoolTensor<Self>,
164        dim: usize,
165        indices: IntTensor<Self>,
166    ) -> BoolTensor<Self> {
167        kernel::select(tensor, dim, indices)
168    }
169
170    fn bool_select_or(
171        tensor: BoolTensor<Self>,
172        dim: usize,
173        indices: IntTensor<Self>,
174        value: BoolTensor<Self>,
175    ) -> BoolTensor<Self> {
176        kernel::select_assign(tensor, dim, indices, value, true)
177    }
178
179    fn bool_flip(tensor: BoolTensor<Self>, axes: &[usize]) -> BoolTensor<Self> {
180        let dtype = tensor.dtype;
181        kernel::flip(tensor, axes, dtype)
182    }
183
184    fn bool_unfold(
185        tensor: FloatTensor<Self>,
186        dim: usize,
187        size: usize,
188        step: usize,
189    ) -> FloatTensor<Self> {
190        unfold(tensor, dim, size, step)
191    }
192
193    fn bool_mask_where(
194        tensor: BoolTensor<Self>,
195        mask: BoolTensor<Self>,
196        value: BoolTensor<Self>,
197    ) -> BoolTensor<Self> {
198        let dtype = tensor.dtype;
199        kernel::mask_where_auto(tensor, mask, value, dtype)
200    }
201
202    fn bool_mask_fill(
203        tensor: BoolTensor<Self>,
204        mask: BoolTensor<Self>,
205        value: Scalar,
206    ) -> BoolTensor<Self> {
207        let dtype = tensor.dtype;
208        kernel::mask_fill_auto(
209            tensor,
210            mask,
211            InputScalar::new(value, dtype_to_storage_type(dtype)),
212            dtype,
213        )
214    }
215
216    fn bool_gather(
217        dim: usize,
218        tensor: BoolTensor<Self>,
219        indices: IntTensor<Self>,
220    ) -> BoolTensor<Self> {
221        kernel::gather(dim, tensor, indices)
222    }
223
224    fn bool_scatter_or(
225        dim: usize,
226        tensor: BoolTensor<Self>,
227        indices: IntTensor<Self>,
228        value: BoolTensor<Self>,
229    ) -> BoolTensor<Self> {
230        kernel::scatter(dim, tensor, indices, value, true)
231    }
232
233    fn bool_equal_elem(lhs: BoolTensor<Self>, rhs: Scalar) -> BoolTensor<Self> {
234        let dtype = lhs.dtype;
235        kernel::equal_elem(
236            lhs,
237            InputScalar::new(rhs, dtype_to_storage_type(dtype)),
238            dtype,
239        )
240    }
241}