Skip to main content

ruda_tensor_device/fusion/
mod.rs

1use 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/// Fusion runtime for JIT runtimes.
158#[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}