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
21fn 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 let all_steps_one = slices.iter().all(|info| info.step == 1);
72
73 if all_steps_one {
74 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 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}