ruda_tensor_device/dispatch/
boolean.rs1use crate::{DeviceBackend, DeviceRuntime, FloatElement, IntElement, element::BoolElement};
2use ruprim::elementwise::binary::numeric::{AndOp, OrOp};
3use ruda_tensor::{
4 ExecutionError, Slice,
5 ops::BoolTensorOps,
6 tensor::{BoolTensor, Device, FloatTensor, IntTensor},
7};
8use ruda_tensor::{Scalar, Shape, TensorData};
9use ruda_core::tensor::{BoolDType, BoolStore, DType, FloatDType, IntDType};
10use ruda_kernel::dsl::prelude::InputScalar;
11use std::ops::Range;
12
13use super::{expand, numeric, permute, unfold};
14
15impl<R, F, I, BT> BoolTensorOps<Self> for DeviceBackend<R, F, I, BT>
16where
17 R: DeviceRuntime,
18 F: FloatElement,
19 I: IntElement,
20 BT: BoolElement,
21{
22 fn bool_empty(shape: Shape, device: &Device<Self>, dtype: BoolDType) -> BoolTensor<Self> {
23 super::empty(shape, device, dtype.into())
24 }
25
26 fn bool_zeros(shape: Shape, device: &Device<Self>, dtype: BoolDType) -> BoolTensor<Self> {
27 numeric::zeros(device.clone(), shape, dtype.into())
28 }
29
30 fn bool_ones(shape: Shape, device: &Device<Self>, dtype: BoolDType) -> BoolTensor<Self> {
31 numeric::ones(device.clone(), shape, dtype.into())
32 }
33
34 async fn bool_into_data(tensor: BoolTensor<Self>) -> Result<TensorData, ExecutionError> {
35 super::into_data(tensor).await
36 }
37
38 fn bool_from_data(data: TensorData, device: &Device<Self>) -> BoolTensor<Self> {
39 if !matches!(
40 data.dtype,
41 DType::Bool(BoolStore::U8) | DType::Bool(BoolStore::U32)
42 ) {
43 unimplemented!("Unsupported dtype for `bool_from_data` {:?}", data.dtype);
44 }
45 super::from_data(data, device)
46 }
47
48 fn bool_into_int(tensor: BoolTensor<Self>, out_dtype: IntDType) -> IntTensor<Self> {
49 ruprim::elementwise::cast::bool_cast(tensor, out_dtype.into())
50 }
51
52 fn bool_device(tensor: &BoolTensor<Self>) -> Device<Self> {
53 tensor.device.clone()
54 }
55
56 fn bool_to_device(tensor: BoolTensor<Self>, device: &Device<Self>) -> BoolTensor<Self> {
57 super::to_device(tensor, device)
58 }
59
60 fn bool_reshape(tensor: BoolTensor<Self>, shape: Shape) -> BoolTensor<Self> {
61 super::reshape(tensor, shape)
62 }
63
64 fn bool_slice(tensor: BoolTensor<Self>, slices: &[Slice]) -> BoolTensor<Self> {
65 let all_steps_one = slices.iter().all(|info| info.step == 1);
67
68 if all_steps_one {
69 let simple_ranges: Vec<Range<usize>> = slices
71 .iter()
72 .enumerate()
73 .map(|(i, slice)| slice.to_range(tensor.meta.shape()[i]))
74 .collect();
75
76 ruprim::indexing::slice(tensor, &simple_ranges)
77 } else {
78 ruprim::indexing::slice_with_steps(tensor, slices)
80 }
81 }
82
83 fn bool_slice_assign(
84 tensor: BoolTensor<Self>,
85 ranges: &[Slice],
86 value: BoolTensor<Self>,
87 ) -> BoolTensor<Self> {
88 ruprim::indexing::slice_assign(tensor, ranges, value)
89 }
90
91 fn bool_equal(lhs: BoolTensor<Self>, rhs: BoolTensor<Self>) -> BoolTensor<Self> {
92 let dtype = lhs.dtype;
93 ruprim::elementwise::comparison::equal(lhs, rhs, dtype)
94 }
95
96 fn bool_not(tensor: BoolTensor<Self>) -> BoolTensor<Self> {
97 let dtype = tensor.dtype;
98 let scalar = match dtype {
99 DType::Bool(BoolStore::U32) => InputScalar::new(u32::false_val(), dtype),
100 DType::Bool(BoolStore::U8) => InputScalar::new(u8::false_val(), dtype),
101 other => unimplemented!("Unsupported dtype for `bool_from_data` {other:?}"),
102 };
103 ruprim::elementwise::comparison::equal_elem(tensor, scalar, dtype)
104 }
105
106 fn bool_and(lhs: BoolTensor<Self>, rhs: BoolTensor<Self>) -> BoolTensor<Self> {
107 ruprim::elementwise::binary::numeric::launch_binop::<R, AndOp>(lhs, rhs)
108 }
109
110 fn bool_or(lhs: BoolTensor<Self>, rhs: BoolTensor<Self>) -> BoolTensor<Self> {
111 ruprim::elementwise::binary::numeric::launch_binop::<R, OrOp>(lhs, rhs)
112 }
113
114 fn bool_into_float(tensor: BoolTensor<Self>, out_dtype: FloatDType) -> FloatTensor<Self> {
115 ruprim::elementwise::cast::bool_cast(tensor, out_dtype.into())
116 }
117
118 fn bool_swap_dims(mut tensor: BoolTensor<Self>, dim1: usize, dim2: usize) -> BoolTensor<Self> {
119 tensor.meta.swap(dim1, dim2);
120
121 tensor
122 }
123
124 fn bool_repeat_dim(tensor: BoolTensor<Self>, dim: usize, times: usize) -> BoolTensor<Self> {
125 ruprim::indexing::repeat_dim(tensor, dim, times)
126 }
127
128 fn bool_permute(tensor: BoolTensor<Self>, axes: &[usize]) -> BoolTensor<Self> {
129 permute(tensor, axes)
130 }
131
132 fn bool_expand(tensor: BoolTensor<Self>, shape: Shape) -> BoolTensor<Self> {
133 expand(tensor, shape)
134 }
135
136 fn bool_select(
137 tensor: BoolTensor<Self>,
138 dim: usize,
139 indices: IntTensor<Self>,
140 ) -> BoolTensor<Self> {
141 ruprim::indexing::select(tensor, dim, indices)
142 }
143
144 fn bool_select_or(
145 tensor: BoolTensor<Self>,
146 dim: usize,
147 indices: IntTensor<Self>,
148 value: BoolTensor<Self>,
149 ) -> BoolTensor<Self> {
150 ruprim::indexing::select_assign(tensor, dim, indices, value, true)
151 }
152
153 fn bool_flip(tensor: BoolTensor<Self>, axes: &[usize]) -> BoolTensor<Self> {
154 let dtype = tensor.dtype;
155 ruprim::indexing::flip(tensor, axes, dtype)
156 }
157
158 fn bool_unfold(
159 tensor: FloatTensor<Self>,
160 dim: usize,
161 size: usize,
162 step: usize,
163 ) -> FloatTensor<Self> {
164 unfold(tensor, dim, size, step)
165 }
166
167 fn bool_mask_where(
168 tensor: BoolTensor<Self>,
169 mask: BoolTensor<Self>,
170 value: BoolTensor<Self>,
171 ) -> BoolTensor<Self> {
172 let dtype = tensor.dtype;
173 ruprim::elementwise::mask::mask_where_auto(tensor, mask, value, dtype)
174 }
175
176 fn bool_mask_fill(
177 tensor: BoolTensor<Self>,
178 mask: BoolTensor<Self>,
179 value: Scalar,
180 ) -> BoolTensor<Self> {
181 let dtype = tensor.dtype;
182 ruprim::elementwise::mask::mask_fill_auto(tensor, mask, InputScalar::new(value, dtype), dtype)
183 }
184
185 fn bool_gather(
186 dim: usize,
187 tensor: BoolTensor<Self>,
188 indices: IntTensor<Self>,
189 ) -> BoolTensor<Self> {
190 ruprim::indexing::gather(dim, tensor, indices)
191 }
192
193 fn bool_scatter_or(
194 dim: usize,
195 tensor: BoolTensor<Self>,
196 indices: IntTensor<Self>,
197 value: BoolTensor<Self>,
198 ) -> BoolTensor<Self> {
199 ruprim::indexing::scatter(dim, tensor, indices, value, true)
200 }
201
202 fn bool_equal_elem(lhs: BoolTensor<Self>, rhs: Scalar) -> BoolTensor<Self> {
203 let dtype = lhs.dtype;
204 ruprim::elementwise::comparison::equal_elem(lhs, InputScalar::new(rhs, dtype), dtype)
205 }
206}