Skip to main content

ruda_tensor_device/dispatch/
boolean.rs

1use 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        // Check if all steps are 1
66        let all_steps_one = slices.iter().all(|info| info.step == 1);
67
68        if all_steps_one {
69            // Use optimized slice for step=1
70            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            // Use slice with steps kernel
79            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}