Skip to main content

burn_cubecl/
fusion.rs

1use crate::{CubeBackend, CubeRuntime, kernel, tensor::CubeTensor};
2use burn_backend::tensor::{BoolTensor, FloatTensor, IntTensor, QuantizedTensor};
3use burn_backend::{DType, Shape};
4pub use burn_cubecl_fusion::{CubeFusionHandle, FallbackOperation};
5use burn_fusion::UnfusedOp;
6use burn_fusion::{
7    FusionBackend, FusionRuntime,
8    stream::{Operation, OrderedExecution},
9};
10use burn_ir::{BackendIr, TensorHandle};
11use burn_std::Metadata;
12use core::marker::PhantomData;
13use std::sync::Arc;
14
15mod registry;
16pub use burn_cubecl_fusion::optim::{CubeOptimization, CubeOptimizationState, FusedOperation};
17pub use registry::{
18    BUILTIN_NAMES, CubeFuser, OptimizationProvider, RegistryError, register, remove,
19};
20
21impl<R> burn_fusion::Optimization<FusionCubeRuntime<R>> for CubeOptimization<R>
22where
23    R: CubeRuntime,
24{
25    fn execute(
26        &mut self,
27        context: &mut burn_fusion::stream::Context<
28            <FusionCubeRuntime<R> as FusionRuntime>::FusionHandle,
29        >,
30        execution: &OrderedExecution<FusionCubeRuntime<R>>,
31    ) {
32        self.run(context, &|index| {
33            let operation = execution.operation_within_optimization(index);
34            Box::new(FallbackOperationWrapper::new(operation))
35        })
36    }
37
38    fn to_state(&self) -> CubeOptimizationState {
39        Self::to_state(self)
40    }
41
42    fn from_state(device: &R::Device, state: CubeOptimizationState) -> Self {
43        registry::restore::<R>(device, state)
44    }
45}
46
47struct FallbackOperationWrapper<O: Clone> {
48    operation: O,
49}
50
51impl<O: Clone> FallbackOperationWrapper<O> {
52    fn new(op: O) -> Self {
53        Self { operation: op }
54    }
55}
56
57impl<R: CubeRuntime> FallbackOperation<R>
58    for FallbackOperationWrapper<Arc<dyn Operation<FusionCubeRuntime<R>>>>
59{
60    fn run(&self, context: &mut burn_fusion::stream::Context<CubeFusionHandle<R>>) {
61        self.operation.as_ref().execute(&mut context.handles);
62    }
63}
64
65impl<R: CubeRuntime> FallbackOperation<R>
66    for FallbackOperationWrapper<UnfusedOp<FusionCubeRuntime<R>>>
67{
68    fn run(&self, context: &mut burn_fusion::stream::Context<CubeFusionHandle<R>>) {
69        self.operation.execute(&mut context.handles);
70    }
71}
72
73impl<R: CubeRuntime> BackendIr for CubeBackend<R> {
74    type Handle = CubeFusionHandle<R>;
75
76    fn float_tensor(handle: TensorHandle<Self::Handle>) -> FloatTensor<Self> {
77        into_tensor(handle.handle, handle.shape)
78    }
79
80    fn int_tensor(handle: TensorHandle<Self::Handle>) -> IntTensor<Self> {
81        into_tensor(handle.handle, handle.shape)
82    }
83
84    fn bool_tensor(handle: TensorHandle<Self::Handle>) -> BoolTensor<Self> {
85        into_tensor(handle.handle, handle.shape)
86    }
87
88    fn quantized_tensor(handle: TensorHandle<Self::Handle>) -> QuantizedTensor<Self> {
89        into_tensor(handle.handle, handle.shape)
90    }
91
92    fn float_tensor_handle(tensor: FloatTensor<Self>) -> Self::Handle {
93        tensor.into()
94    }
95
96    fn int_tensor_handle(tensor: IntTensor<Self>) -> Self::Handle {
97        tensor.into()
98    }
99
100    fn bool_tensor_handle(tensor: BoolTensor<Self>) -> Self::Handle {
101        tensor.into()
102    }
103
104    fn quantized_tensor_handle(tensor: QuantizedTensor<Self>) -> Self::Handle {
105        tensor.into()
106    }
107}
108
109impl<R: CubeRuntime> FusionRuntime for FusionCubeRuntime<R> {
110    type OptimizationState = CubeOptimizationState;
111    type Optimization = CubeOptimization<R>;
112    type FusionHandle = CubeFusionHandle<R>;
113    type FusionDevice = R::CubeDevice;
114
115    fn fusers(device: R::Device) -> Vec<Box<dyn burn_fusion::OperationFuser<Self::Optimization>>> {
116        registry::fusers::<R>(&device)
117    }
118}
119
120/// Fusion runtime for JIT runtimes.
121#[derive(Debug)]
122pub struct FusionCubeRuntime<R: CubeRuntime> {
123    _b: PhantomData<R>,
124}
125
126impl<R: CubeRuntime> FusionBackend for CubeBackend<R> {
127    type FusionRuntime = FusionCubeRuntime<R>;
128
129    type FullPrecisionBackend = CubeBackend<R>;
130
131    fn cast_float(tensor: FloatTensor<Self>, dtype: DType) -> Self::Handle {
132        kernel::cast(tensor, dtype).into()
133    }
134
135    fn memory_persistent(device: &Self::Device, enabled: bool) {
136        use cubecl::MemoryAllocationMode;
137
138        let client = R::client(device);
139        let mode = match enabled {
140            true => MemoryAllocationMode::Persistent,
141            false => MemoryAllocationMode::Auto,
142        };
143        // Safety: called from the fusion execution thread, whose stream is the
144        // one every fused operation allocates on.
145        unsafe { client.allocation_mode(mode) };
146    }
147}
148
149fn into_tensor<R: CubeRuntime>(handle: CubeFusionHandle<R>, shape: Shape) -> CubeTensor<R> {
150    CubeTensor {
151        client: handle.client.clone(),
152        handle: handle.handle.clone(),
153        device: handle.device.clone(),
154        meta: Box::new(Metadata::new(shape, handle.strides.clone())),
155        dtype: handle.dtype,
156        qparams: handle.qparams.clone(),
157    }
158}
159
160impl<R: CubeRuntime> From<CubeTensor<R>> for CubeFusionHandle<R> {
161    fn from(value: CubeTensor<R>) -> Self {
162        Self {
163            client: value.client.clone(),
164            handle: value.handle.clone(),
165            device: value.device.clone(),
166            strides: value.meta.strides.clone(),
167            dtype: value.dtype,
168            qparams: value.qparams.clone(),
169        }
170    }
171}