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#[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}