ruda_tensor_device/fusion/
mod.rs1use crate::BoolElement;
2use crate::{DeviceBackend, DeviceRuntime, FloatElement, IntElement, RudaTensor};
3use ruda_tensor::tensor::{BoolTensor, FloatTensor, IntTensor, QuantizedTensor};
4use ruda_tensor::{DType, Shape, quantization::QuantScheme};
5use ruda_fusion::device::optim::reduce::ReduceSettings;
6use ruda_fusion::device::optim::reduce_broadcasted::ReduceBroadcastedFuser;
7use ruda_fusion::device::{
8 RudaFusionHandle, FallbackOperation,
9 optim::{
10 RudaOptimization, RudaOptimizationState,
11 elemwise::{ElementWiseFuser, ElemwiseOptimization},
12 matmul::{MatmulFuser, MatmulOptimization},
13 reduce::{ReduceFuser, ReduceOptimization},
14 reduce_broadcasted::ReduceBroadcastedOptimization,
15 },
16};
17use ruda_fusion::UnfusedOp;
18use ruda_fusion::{
19 FusionBackend, FusionRuntime,
20 stream::{Operation, OrderedExecution},
21};
22use ruda_tensor::graph::{BackendIr, TensorHandle};
23use ruda_fusion::device::tensor::into_tensor;
24use core::marker::PhantomData;
25use std::sync::Arc;
26
27impl<R> ruda_fusion::Optimization<DeviceFusionRuntime<R>> for RudaOptimization<R>
28where
29 R: DeviceRuntime,
30{
31 fn execute(
32 &mut self,
33 context: &mut ruda_fusion::stream::Context<
34 <DeviceFusionRuntime<R> as FusionRuntime>::FusionHandle,
35 >,
36 execution: &OrderedExecution<DeviceFusionRuntime<R>>,
37 ) {
38 match self {
39 Self::ElementWise(op) => op.execute(context),
40 Self::Matmul(op) => op.execute(context, |index| {
41 let operation = execution.operation_within_optimization(index);
42 Box::new(FallbackOperationWrapper::new(operation))
43 }),
44 Self::Reduce(op) => op.execute(context, |index| {
45 let operation = execution.operation_within_optimization(index);
46 Box::new(FallbackOperationWrapper::new(operation))
47 }),
48 Self::ReduceBroadcasted(op) => op.execute(context, |index| {
49 let operation = execution.operation_within_optimization(index);
50 Box::new(FallbackOperationWrapper::new(operation))
51 }),
52 }
53 }
54
55 fn to_state(&self) -> RudaOptimizationState {
56 self.to_opt_state()
57 }
58
59 fn from_state(device: &R::Device, state: RudaOptimizationState) -> Self {
60 match state {
61 RudaOptimizationState::ElementWise(state) => {
62 Self::ElementWise(ElemwiseOptimization::from_state(device, state))
63 }
64 RudaOptimizationState::Matmul(state) => {
65 Self::Matmul(MatmulOptimization::from_state(device, state))
66 }
67 RudaOptimizationState::Reduce(state) => {
68 Self::Reduce(ReduceOptimization::from_state(device, state))
69 }
70 RudaOptimizationState::ReduceBroadcasted(state) => {
71 Self::ReduceBroadcasted(ReduceBroadcastedOptimization::from_state(device, state))
72 }
73 }
74 }
75}
76
77struct FallbackOperationWrapper<O: Clone> {
78 operation: O,
79}
80
81impl<O: Clone> FallbackOperationWrapper<O> {
82 fn new(op: O) -> Self {
83 Self { operation: op }
84 }
85}
86
87impl<R: DeviceRuntime> FallbackOperation<R>
88 for FallbackOperationWrapper<Arc<dyn Operation<DeviceFusionRuntime<R>>>>
89{
90 fn run(&self, context: &mut ruda_fusion::stream::Context<RudaFusionHandle<R>>) {
91 self.operation.as_ref().execute(&mut context.handles);
92 }
93}
94
95impl<R: DeviceRuntime> FallbackOperation<R>
96 for FallbackOperationWrapper<UnfusedOp<DeviceFusionRuntime<R>>>
97{
98 fn run(&self, context: &mut ruda_fusion::stream::Context<RudaFusionHandle<R>>) {
99 self.operation.execute(&mut context.handles);
100 }
101}
102
103impl<R: DeviceRuntime, F: FloatElement, I: IntElement, BT: BoolElement> BackendIr
104 for DeviceBackend<R, F, I, BT>
105{
106 type Handle = RudaFusionHandle<R>;
107
108 fn float_tensor(handle: TensorHandle<Self::Handle>) -> FloatTensor<Self> {
109 into_tensor(handle.handle, handle.shape)
110 }
111
112 fn int_tensor(handle: TensorHandle<Self::Handle>) -> IntTensor<Self> {
113 into_tensor(handle.handle, handle.shape)
114 }
115
116 fn bool_tensor(handle: TensorHandle<Self::Handle>) -> BoolTensor<Self> {
117 into_tensor(handle.handle, handle.shape)
118 }
119
120 fn quantized_tensor(handle: TensorHandle<Self::Handle>) -> QuantizedTensor<Self> {
121 into_tensor(handle.handle, handle.shape)
122 }
123
124 fn float_tensor_handle(tensor: FloatTensor<Self>) -> Self::Handle {
125 tensor.into()
126 }
127
128 fn int_tensor_handle(tensor: IntTensor<Self>) -> Self::Handle {
129 tensor.into()
130 }
131
132 fn bool_tensor_handle(tensor: BoolTensor<Self>) -> Self::Handle {
133 tensor.into()
134 }
135
136 fn quantized_tensor_handle(tensor: QuantizedTensor<Self>) -> Self::Handle {
137 tensor.into()
138 }
139}
140
141impl<R: DeviceRuntime> FusionRuntime for DeviceFusionRuntime<R> {
142 type OptimizationState = RudaOptimizationState;
143 type Optimization = RudaOptimization<R>;
144 type FusionHandle = RudaFusionHandle<R>;
145 type FusionDevice = R::RudaDevice;
146
147 fn fusers(device: R::Device) -> Vec<Box<dyn ruda_fusion::OperationFuser<Self::Optimization>>> {
148 vec![
149 Box::new(ElementWiseFuser::new(device.clone())),
150 Box::new(MatmulFuser::new(device.clone())),
151 Box::new(ReduceFuser::new(device.clone(), ReduceSettings::Always)),
152 Box::new(ReduceBroadcastedFuser::new(device.clone())),
153 ]
154 }
155}
156
157#[derive(Debug)]
159pub struct DeviceFusionRuntime<R: DeviceRuntime> {
160 _b: PhantomData<R>,
161}
162
163impl<R: DeviceRuntime, F: FloatElement, I: IntElement, BT: BoolElement> FusionBackend
164 for DeviceBackend<R, F, I, BT>
165{
166 type FusionRuntime = DeviceFusionRuntime<R>;
167
168 type FullPrecisionBackend = DeviceBackend<R, f32, i32, BT>;
169
170 fn avg_pool3d_output_size(input: [usize; 3], kernel: [usize; 3], stride: [usize; 3],
171 padding: [usize; 3], ceil: bool) -> [usize; 3] {
172 rudnn::pooling::avg_pool3d_output_size(input, kernel, stride, padding, ceil)
173 }
174
175 fn cast_float(tensor: FloatTensor<Self>, dtype: DType) -> Self::Handle {
176 ruprim::elementwise::cast::cast(tensor, dtype).into()
177 }
178
179 fn q_swap_dims_scheme(
180 scheme: QuantScheme,
181 rank: usize,
182 dim1: usize,
183 dim2: usize,
184 ) -> QuantScheme {
185 ruda_kernel::tensor::permutation::swap_dims_scheme(scheme, rank, dim1, dim2)
186 }
187
188 fn q_permute_scheme(scheme: QuantScheme, axes: &[usize]) -> QuantScheme {
189 ruda_kernel::tensor::permutation::permute_scheme(scheme, axes)
190 }
191}