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