Skip to main content

burn_cubecl/ops/
module.rs

1use crate::{
2    CubeBackend, CubeRuntime,
3    kernel::{self, conv::ConvTranspose2dStrategy},
4};
5use burn_backend::tensor::{BoolTensor, FloatTensor, IntTensor};
6use burn_backend::{
7    TensorMetadata,
8    ops::{
9        AttentionModuleOptions, ConvOptions, ConvTransposeOptions, DeformConv2dBackward,
10        DeformConvOptions, InterpolateOptions, MaxPool2dBackward, MaxPool2dWithIndices, ModuleOps,
11    },
12};
13use burn_std::IntDType;
14
15impl<R> ModuleOps<Self> for CubeBackend<R>
16where
17    R: CubeRuntime,
18{
19    fn conv1d(
20        x: FloatTensor<Self>,
21        weight: FloatTensor<Self>,
22        bias: Option<FloatTensor<Self>>,
23        options: ConvOptions<1>,
24    ) -> FloatTensor<Self> {
25        kernel::conv::conv_forward::<R, 1>(x, weight, bias, options, Default::default()).unwrap()
26    }
27
28    fn conv1d_x_backward(
29        x: FloatTensor<Self>,
30        weight: FloatTensor<Self>,
31        output_grad: FloatTensor<Self>,
32        options: ConvOptions<1>,
33    ) -> FloatTensor<Self> {
34        kernel::conv::conv_data_backward(
35            output_grad,
36            weight,
37            x.shape(),
38            options,
39            Default::default(),
40        )
41        .unwrap()
42    }
43
44    fn conv1d_weight_backward(
45        x: FloatTensor<Self>,
46        weight: FloatTensor<Self>,
47        output_grad: FloatTensor<Self>,
48        options: ConvOptions<1>,
49    ) -> FloatTensor<Self> {
50        kernel::conv::conv_weight_backward::<R, 1>(
51            x,
52            output_grad,
53            weight.shape(),
54            options,
55            Default::default(),
56        )
57        .unwrap()
58    }
59
60    fn conv2d(
61        x: FloatTensor<Self>,
62        weight: FloatTensor<Self>,
63        bias: Option<FloatTensor<Self>>,
64        options: ConvOptions<2>,
65    ) -> FloatTensor<Self> {
66        kernel::conv::conv_forward::<R, 2>(x, weight, bias, options, Default::default()).unwrap()
67    }
68
69    fn conv2d_x_backward(
70        x: FloatTensor<Self>,
71        weight: FloatTensor<Self>,
72        output_grad: FloatTensor<Self>,
73        options: ConvOptions<2>,
74    ) -> FloatTensor<Self> {
75        kernel::conv::conv_data_backward(
76            output_grad,
77            weight,
78            x.shape(),
79            options,
80            Default::default(),
81        )
82        .unwrap()
83    }
84
85    fn conv2d_weight_backward(
86        x: FloatTensor<Self>,
87        weight: FloatTensor<Self>,
88        output_grad: FloatTensor<Self>,
89        options: ConvOptions<2>,
90    ) -> FloatTensor<Self> {
91        kernel::conv::conv_weight_backward::<R, 2>(
92            x,
93            output_grad,
94            weight.shape(),
95            options,
96            Default::default(),
97        )
98        .unwrap()
99    }
100
101    fn deform_conv2d(
102        x: FloatTensor<Self>,
103        offset: FloatTensor<Self>,
104        weight: FloatTensor<Self>,
105        mask: Option<FloatTensor<Self>>,
106        bias: Option<FloatTensor<Self>>,
107        options: DeformConvOptions<2>,
108    ) -> FloatTensor<Self> {
109        kernel::conv::deform_conv2d(x, offset, weight, mask, bias, options).unwrap()
110    }
111
112    fn deform_conv2d_backward(
113        x: FloatTensor<Self>,
114        offset: FloatTensor<Self>,
115        weight: FloatTensor<Self>,
116        mask: Option<FloatTensor<Self>>,
117        bias: Option<FloatTensor<Self>>,
118        output_grad: FloatTensor<Self>,
119        options: DeformConvOptions<2>,
120    ) -> DeformConv2dBackward<Self> {
121        let (x, o, w, m, b) = kernel::conv::deform_conv2d_backward(
122            x,
123            offset,
124            weight,
125            mask,
126            bias,
127            output_grad,
128            options,
129        )
130        .unwrap();
131        DeformConv2dBackward::new(x, o, w, m, b)
132    }
133
134    fn conv3d(
135        x: FloatTensor<Self>,
136        weight: FloatTensor<Self>,
137        bias: Option<FloatTensor<Self>>,
138        options: ConvOptions<3>,
139    ) -> FloatTensor<Self> {
140        kernel::conv::conv_forward::<R, 3>(x, weight, bias, options, Default::default()).unwrap()
141    }
142
143    fn conv3d_x_backward(
144        x: FloatTensor<Self>,
145        weight: FloatTensor<Self>,
146        output_grad: FloatTensor<Self>,
147        options: ConvOptions<3>,
148    ) -> FloatTensor<Self> {
149        kernel::conv::conv_data_backward(
150            output_grad,
151            weight,
152            x.shape(),
153            options,
154            Default::default(),
155        )
156        .unwrap()
157    }
158
159    fn conv3d_weight_backward(
160        x: FloatTensor<Self>,
161        weight: FloatTensor<Self>,
162        output_grad: FloatTensor<Self>,
163        options: ConvOptions<3>,
164    ) -> FloatTensor<Self> {
165        kernel::conv::conv_weight_backward::<R, 3>(
166            x,
167            output_grad,
168            weight.shape(),
169            options,
170            Default::default(),
171        )
172        .unwrap()
173    }
174
175    fn conv_transpose2d(
176        x: FloatTensor<Self>,
177        weight: FloatTensor<Self>,
178        bias: Option<FloatTensor<Self>>,
179        options: ConvTransposeOptions<2>,
180    ) -> FloatTensor<Self> {
181        kernel::conv::conv_transpose2d(x, weight, bias, options, ConvTranspose2dStrategy::default())
182            .unwrap()
183    }
184
185    fn conv_transpose3d(
186        x: FloatTensor<Self>,
187        weight: FloatTensor<Self>,
188        bias: Option<FloatTensor<Self>>,
189        options: ConvTransposeOptions<3>,
190    ) -> FloatTensor<Self> {
191        kernel::conv::conv_transpose3d(x, weight, bias, options).expect("Kernel to never fail")
192    }
193
194    fn avg_pool2d(
195        x: FloatTensor<Self>,
196        kernel_size: [usize; 2],
197        stride: [usize; 2],
198        padding: [usize; 2],
199        count_include_pad: bool,
200        ceil_mode: bool,
201    ) -> FloatTensor<Self> {
202        kernel::pool::avg_pool2d(
203            x,
204            kernel_size,
205            stride,
206            padding,
207            count_include_pad,
208            ceil_mode,
209        )
210    }
211
212    fn avg_pool2d_backward(
213        x: FloatTensor<Self>,
214        grad: FloatTensor<Self>,
215        kernel_size: [usize; 2],
216        stride: [usize; 2],
217        padding: [usize; 2],
218        count_include_pad: bool,
219        ceil_mode: bool,
220    ) -> FloatTensor<Self> {
221        kernel::pool::avg_pool2d_backward(
222            x,
223            grad,
224            kernel_size,
225            stride,
226            padding,
227            count_include_pad,
228            ceil_mode,
229        )
230    }
231
232    fn max_pool2d(
233        x: FloatTensor<Self>,
234        kernel_size: [usize; 2],
235        stride: [usize; 2],
236        padding: [usize; 2],
237        dilation: [usize; 2],
238        ceil_mode: bool,
239    ) -> FloatTensor<Self> {
240        kernel::pool::max_pool2d(x, kernel_size, stride, padding, dilation, ceil_mode)
241    }
242
243    fn max_pool2d_with_indices(
244        x: FloatTensor<Self>,
245        kernel_size: [usize; 2],
246        stride: [usize; 2],
247        padding: [usize; 2],
248        dilation: [usize; 2],
249        ceil_mode: bool,
250        indices_dtype: IntDType,
251    ) -> MaxPool2dWithIndices<Self> {
252        let (output, indices) = kernel::pool::max_pool2d_with_indices(
253            x,
254            kernel_size,
255            stride,
256            padding,
257            dilation,
258            ceil_mode,
259            indices_dtype.into(),
260        );
261
262        MaxPool2dWithIndices::new(output, indices)
263    }
264
265    fn max_pool2d_with_indices_backward(
266        x: FloatTensor<Self>,
267        kernel_size: [usize; 2],
268        stride: [usize; 2],
269        padding: [usize; 2],
270        dilation: [usize; 2],
271        ceil_mode: bool,
272        output_grad: FloatTensor<Self>,
273        indices: IntTensor<Self>,
274    ) -> MaxPool2dBackward<Self> {
275        MaxPool2dBackward::new(kernel::pool::max_pool2d_with_indices_backward(
276            x,
277            output_grad,
278            indices,
279            kernel_size,
280            stride,
281            padding,
282            dilation,
283            ceil_mode,
284        ))
285    }
286
287    fn adaptive_avg_pool2d(x: FloatTensor<Self>, output_size: [usize; 2]) -> FloatTensor<Self> {
288        kernel::pool::adaptive_avg_pool2d(x, output_size)
289    }
290
291    fn adaptive_avg_pool2d_backward(
292        x: FloatTensor<Self>,
293        grad: FloatTensor<Self>,
294    ) -> FloatTensor<Self> {
295        kernel::pool::adaptive_avg_pool2d_backward(x, grad)
296    }
297
298    fn adaptive_avg_pool3d(_x: FloatTensor<Self>, _output_size: [usize; 3]) -> FloatTensor<Self> {
299        todo!("CubeCL backend does not yet support adaptive_avg_pool3d.")
300    }
301
302    fn adaptive_avg_pool3d_backward(
303        _x: FloatTensor<Self>,
304        _grad: FloatTensor<Self>,
305    ) -> FloatTensor<Self> {
306        todo!("CubeCL backend does not yet support adaptive_avg_pool3d_backward.")
307    }
308
309    fn interpolate(
310        x: FloatTensor<Self>,
311        output_size: [usize; 2],
312        options: InterpolateOptions,
313    ) -> FloatTensor<Self> {
314        kernel::interpolate::interpolate(x, output_size, options, Default::default()).unwrap()
315    }
316
317    fn interpolate_backward(
318        x: FloatTensor<Self>,
319        grad: FloatTensor<Self>,
320        output_size: [usize; 2],
321        options: InterpolateOptions,
322    ) -> FloatTensor<Self> {
323        kernel::interpolate::interpolate_backward(x, grad, output_size, options)
324    }
325
326    fn attention(
327        query: FloatTensor<Self>,
328        key: FloatTensor<Self>,
329        value: FloatTensor<Self>,
330        mask: Option<BoolTensor<Self>>,
331        attn_bias: Option<FloatTensor<Self>>,
332        options: AttentionModuleOptions,
333    ) -> FloatTensor<Self> {
334        // Fall back to naive attention for features the flash kernel doesn't support.
335        if attn_bias.is_some() || options.softcap.is_some() || options.scale.is_some() {
336            return burn_backend::ops::attention::attention_fallback::<Self>(
337                query, key, value, mask, attn_bias, options,
338            );
339        }
340
341        kernel::attention::attention(
342            query,
343            key,
344            value,
345            mask,
346            attn_bias,
347            options,
348            Default::default(),
349        )
350        .expect("Kernel to never fail")
351    }
352
353    fn has_ctc_loss_backward() -> bool {
354        true
355    }
356
357    fn ctc_loss(
358        log_probs: FloatTensor<Self>,
359        targets: IntTensor<Self>,
360        input_lengths: IntTensor<Self>,
361        target_lengths: IntTensor<Self>,
362        blank: usize,
363    ) -> FloatTensor<Self> {
364        kernel::ctc::ctc_loss(log_probs, targets, input_lengths, target_lengths, blank)
365    }
366
367    fn ctc_loss_backward(
368        log_probs: FloatTensor<Self>,
369        targets: IntTensor<Self>,
370        input_lengths: IntTensor<Self>,
371        target_lengths: IntTensor<Self>,
372        grad_loss: FloatTensor<Self>,
373        blank: usize,
374    ) -> FloatTensor<Self> {
375        let (log_alpha_full, log_beta_full, nll) = kernel::ctc::ctc_alpha_beta(
376            log_probs.clone(),
377            targets.clone(),
378            input_lengths.clone(),
379            target_lengths,
380            blank,
381        );
382        burn_backend::ops::ctc::ctc_grad_from_alpha_beta_default::<Self>(
383            log_probs,
384            targets,
385            input_lengths,
386            grad_loss,
387            log_alpha_full,
388            log_beta_full,
389            nll,
390            blank,
391        )
392    }
393
394    fn rfft(
395        signal: FloatTensor<Self>,
396        dim: usize,
397        n: Option<usize>,
398    ) -> (FloatTensor<Self>, FloatTensor<Self>) {
399        kernel::fft::rfft(signal, dim, n)
400    }
401
402    fn irfft(
403        spectrum_re: FloatTensor<Self>,
404        spectrum_im: FloatTensor<Self>,
405        dim: usize,
406        n: Option<usize>,
407    ) -> FloatTensor<Self> {
408        kernel::fft::irfft(spectrum_re, spectrum_im, dim, n)
409    }
410}