Skip to main content

ruda_tensor_device/
backend.rs

1use crate::{DeviceRuntime, FloatElement, IntElement, element::BoolElement, RudaTensor};
2use ruda_tensor::{
3    Backend, BackendTypes, DTypeUsage, DTypeUsageSet, DeviceOps, ExecutionError, TensorData,
4};
5use ruda_core::tensor::DType;
6use ruda_core::ir::features::{MmaConfig, TypeUsage};
7use ruda::runtime::server::ComputeServer;
8use std::marker::PhantomData;
9
10#[cfg(not(feature = "fusion"))]
11use ruda_tensor::tensor::{BoolTensor, FloatTensor, IntTensor, QuantizedTensor};
12#[cfg(not(feature = "fusion"))]
13use ruda_tensor::graph::{BackendIr, TensorHandle};
14
15/// Generic tensor backend that can be compiled just-in-time to any shader runtime
16#[derive(new)]
17pub struct DeviceBackend<R: DeviceRuntime, F: FloatElement, I: IntElement, BT: BoolElement> {
18    _runtime: PhantomData<R>,
19    _float_elem: PhantomData<F>,
20    _int_elem: PhantomData<I>,
21    _bool_elem: PhantomData<BT>,
22}
23
24impl<R, F, I, BT> BackendTypes for DeviceBackend<R, F, I, BT>
25where
26    R: DeviceRuntime,
27    R::Server: ComputeServer,
28    R::Device: DeviceOps,
29    F: FloatElement,
30    I: IntElement,
31    BT: BoolElement,
32{
33    type Device = R::Device;
34
35    type FloatElem = F;
36    type IntElem = I;
37    type BoolElem = BT;
38
39    type FloatTensorPrimitive = RudaTensor<R>;
40    type IntTensorPrimitive = RudaTensor<R>;
41    type BoolTensorPrimitive = RudaTensor<R>;
42    type QuantizedTensorPrimitive = RudaTensor<R>;
43}
44
45impl<R, F, I, BT> Backend for DeviceBackend<R, F, I, BT>
46where
47    R: DeviceRuntime,
48    R::Server: ComputeServer,
49    R::Device: DeviceOps,
50    F: FloatElement,
51    I: IntElement,
52    BT: BoolElement,
53{
54    fn name(device: &Self::Device) -> String {
55        let client = R::client(device);
56        format!("ruda<{}>", R::name(&client))
57    }
58
59    fn seed(_device: &Self::Device, seed: u64) {
60        rurand::seed(seed);
61    }
62
63    fn ad_enabled(_device: &Self::Device) -> bool {
64        false
65    }
66
67    fn sync(device: &Self::Device) -> Result<(), ExecutionError> {
68        let client = R::client(device);
69        futures_lite::future::block_on(client.sync()).map_err(|err| ExecutionError::WithContext {
70            reason: format!("{err}"),
71        })
72    }
73
74    fn memory_persistent_allocations<
75        Output: Send,
76        Input: Send,
77        Func: Fn(Input) -> Output + Send,
78    >(
79        device: &Self::Device,
80        input: Input,
81        func: Func,
82    ) -> Output {
83        let client = R::client(device);
84        client.memory_persistent_allocation(input, func).unwrap()
85    }
86
87    fn memory_cleanup(device: &Self::Device) {
88        let client = R::client(device);
89        client.memory_cleanup();
90    }
91
92    fn staging<'a, Iter>(data: Iter, device: &Self::Device)
93    where
94        Iter: Iterator<Item = &'a mut TensorData>,
95    {
96        let client = R::client(device);
97        client.staging(data.map(|td| &mut td.bytes), false);
98    }
99
100    fn supports_dtype(device: &Self::Device, dtype: DType) -> bool {
101        ruda_kernel::tensor::capability::supports_dtype::<R>(device, dtype)
102    }
103
104    fn dtype_usage(device: &Self::Device, dtype: DType) -> DTypeUsageSet {
105        let client = R::client(device);
106
107        let props = client.properties();
108        let storage = dtype.into();
109        let usage = props.type_usage(storage);
110
111        let mut out = DTypeUsageSet::new();
112
113        if usage.is_superset(TypeUsage::Buffer | TypeUsage::Conversion) {
114            out |= DTypeUsage::Storage;
115        }
116
117        if usage.contains(TypeUsage::Arithmetic) {
118            out |= DTypeUsage::Arithmetic;
119        }
120
121        let has_mma = |cfg: &MmaConfig| {
122            cfg.a_type == storage || cfg.b_type == storage || cfg.cd_type == storage
123        };
124        if props.features.matmul.cmma.iter().any(has_mma)
125            || props.features.matmul.mma.iter().any(has_mma)
126        {
127            out |= DTypeUsage::Accelerated;
128        }
129
130        out
131    }
132
133    fn device_count(type_id: u16) -> usize {
134        let client = R::client(&Default::default());
135        client.device_count(type_id)
136    }
137}
138
139impl<R: DeviceRuntime, F: FloatElement, I: IntElement, BT: BoolElement> core::fmt::Debug
140    for DeviceBackend<R, F, I, BT>
141{
142    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
143        f.write_str("RudaBackend")
144    }
145}
146
147impl<R: DeviceRuntime, F: FloatElement, I: IntElement, BT: BoolElement> Clone
148    for DeviceBackend<R, F, I, BT>
149{
150    fn clone(&self) -> Self {
151        Self::new()
152    }
153}
154
155impl<R: DeviceRuntime, F: FloatElement, I: IntElement, BT: BoolElement> Default
156    for DeviceBackend<R, F, I, BT>
157{
158    fn default() -> Self {
159        Self::new()
160    }
161}
162
163#[cfg(not(feature = "fusion"))]
164impl<R: DeviceRuntime, F: FloatElement, I: IntElement, BT: BoolElement> BackendIr
165    for DeviceBackend<R, F, I, BT>
166{
167    type Handle = RudaTensor<R>;
168
169    fn float_tensor(handle: TensorHandle<Self::Handle>) -> FloatTensor<Self> {
170        handle.handle
171    }
172
173    fn int_tensor(handle: TensorHandle<Self::Handle>) -> IntTensor<Self> {
174        handle.handle
175    }
176
177    fn bool_tensor(handle: TensorHandle<Self::Handle>) -> BoolTensor<Self> {
178        handle.handle
179    }
180
181    fn quantized_tensor(handle: TensorHandle<Self::Handle>) -> QuantizedTensor<Self> {
182        handle.handle
183    }
184
185    fn float_tensor_handle(tensor: FloatTensor<Self>) -> Self::Handle {
186        tensor
187    }
188
189    fn int_tensor_handle(tensor: IntTensor<Self>) -> Self::Handle {
190        tensor
191    }
192
193    fn bool_tensor_handle(tensor: BoolTensor<Self>) -> Self::Handle {
194        tensor
195    }
196
197    fn quantized_tensor_handle(tensor: QuantizedTensor<Self>) -> Self::Handle {
198        tensor
199    }
200}