1use crate::{DeviceBackend, DeviceRuntime, FloatElement, IntElement, element::BoolElement};
2use rudnn::convolution::tensor::ConvTranspose2dStrategy;
3use ruda_tensor::tensor::{BoolTensor, FloatTensor, IntTensor};
4use ruda_tensor::{
5 TensorMetadata,
6 ops::{
7 AttentionModuleOptions, ConvOptions, ConvTransposeOptions, DeformConv2dBackward,
8 DeformConvOptions, InterpolateOptions, MaxPool2dBackward, MaxPool2dWithIndices, ModuleOps,
9 },
10};
11
12fn norm_buffer<R: DeviceRuntime>(tensor: crate::RudaTensor<R>) -> ruda::runtime::normalization::TensorBuffer {
13 ruda::runtime::normalization::TensorBuffer {
14 shape: tensor.meta.shape().clone(), strides: tensor.meta.strides().clone(),
15 handle: tensor.handle, dtype: tensor.dtype,
16 }
17}
18
19fn norm_tensor<R: DeviceRuntime>(
20 buffer: ruda::runtime::normalization::TensorBuffer,
21 client: ruda::runtime::client::ComputeClient<R>, device: R::Device,
22) -> crate::RudaTensor<R> {
23 crate::RudaTensor::new(client, buffer.handle,
24 ruda_core::tensor::Metadata::new(buffer.shape, buffer.strides), device, buffer.dtype)
25}
26
27impl<R, F, I, BT> ModuleOps<Self> for DeviceBackend<R, F, I, BT>
28where
29 R: DeviceRuntime,
30 F: FloatElement,
31 I: IntElement,
32 BT: BoolElement,
33{
34 fn has_layer_norm_backward() -> bool { R::has_native_layer_norm() }
35
36 fn layer_norm(
37 tensor: FloatTensor<Self>, gamma: FloatTensor<Self>,
38 beta: Option<FloatTensor<Self>>, epsilon: f64,
39 ) -> FloatTensor<Self> {
40 if R::has_native_layer_norm() {
41 Self::layer_norm_with_stats(tensor, gamma, beta, epsilon).output
42 } else {
43 Self::layer_norm_default(tensor, gamma, beta, epsilon)
44 }
45 }
46
47 fn layer_norm_with_stats(
48 tensor: FloatTensor<Self>, gamma: FloatTensor<Self>,
49 beta: Option<FloatTensor<Self>>, epsilon: f64,
50 ) -> ruda_tensor::ops::LayerNormOutput<Self> {
51 let client = tensor.client.clone();
52 let device = tensor.device.clone();
53 for other in core::iter::once(&gamma).chain(beta.iter()) {
54 assert_eq!(device, other.device, "LayerNorm device mismatch");
55 assert!(client.same_execution_queue(&other.client), "LayerNorm queue mismatch");
56 }
57 let [output, mean, rstd] = R::layer_norm(
58 &client, norm_buffer(tensor), norm_buffer(gamma), beta.map(norm_buffer), epsilon,
59 );
60 let from = |buffer| norm_tensor(buffer, client.clone(), device.clone());
61 ruda_tensor::ops::LayerNormOutput { output: from(output), mean: from(mean), rstd: from(rstd) }
62 }
63
64 fn layer_norm_backward(
65 tensor: FloatTensor<Self>, gamma: FloatTensor<Self>, grad: FloatTensor<Self>,
66 mean: FloatTensor<Self>, rstd: FloatTensor<Self>,
67 ) -> ruda_tensor::ops::LayerNormBackward<Self> {
68 let client = tensor.client.clone();
69 let device = tensor.device.clone();
70 for other in [&gamma, &grad, &mean, &rstd] {
71 assert_eq!(device, other.device, "LayerNorm backward device mismatch");
72 assert!(client.same_execution_queue(&other.client), "LayerNorm backward queue mismatch");
73 }
74 let [input, weight, bias] = R::layer_norm_backward(
75 &client, norm_buffer(tensor), norm_buffer(gamma), norm_buffer(grad),
76 norm_buffer(mean), norm_buffer(rstd),
77 );
78 let from = |buffer| norm_tensor(buffer, client.clone(), device.clone());
79 ruda_tensor::ops::LayerNormBackward { input: from(input), weight: from(weight), bias: from(bias) }
80 }
81
82 fn conv1d(
83 x: FloatTensor<Self>,
84 weight: FloatTensor<Self>,
85 bias: Option<FloatTensor<Self>>,
86 options: ConvOptions<1>,
87 ) -> FloatTensor<Self> {
88 rudnn::convolution::tensor::conv_forward::<R, 1>(x, weight, bias, options, Default::default()).unwrap()
89 }
90
91 fn conv1d_x_backward(
92 x: FloatTensor<Self>,
93 weight: FloatTensor<Self>,
94 output_grad: FloatTensor<Self>,
95 options: ConvOptions<1>,
96 ) -> FloatTensor<Self> {
97 rudnn::convolution::tensor::conv_data_backward(
98 output_grad,
99 weight,
100 x.shape(),
101 options,
102 Default::default(),
103 )
104 .unwrap()
105 }
106
107 fn conv1d_weight_backward(
108 x: FloatTensor<Self>,
109 weight: FloatTensor<Self>,
110 output_grad: FloatTensor<Self>,
111 options: ConvOptions<1>,
112 ) -> FloatTensor<Self> {
113 rudnn::convolution::tensor::conv_weight_backward::<R, 1>(
114 x,
115 output_grad,
116 weight.shape(),
117 options,
118 Default::default(),
119 )
120 .unwrap()
121 }
122
123 fn conv2d(
124 x: FloatTensor<Self>,
125 weight: FloatTensor<Self>,
126 bias: Option<FloatTensor<Self>>,
127 options: ConvOptions<2>,
128 ) -> FloatTensor<Self> {
129 rudnn::convolution::tensor::conv_forward::<R, 2>(x, weight, bias, options, Default::default()).unwrap()
130 }
131
132 fn conv2d_x_backward(
133 x: FloatTensor<Self>,
134 weight: FloatTensor<Self>,
135 output_grad: FloatTensor<Self>,
136 options: ConvOptions<2>,
137 ) -> FloatTensor<Self> {
138 rudnn::convolution::tensor::conv_data_backward(
139 output_grad,
140 weight,
141 x.shape(),
142 options,
143 Default::default(),
144 )
145 .unwrap()
146 }
147
148 fn conv2d_weight_backward(
149 x: FloatTensor<Self>,
150 weight: FloatTensor<Self>,
151 output_grad: FloatTensor<Self>,
152 options: ConvOptions<2>,
153 ) -> FloatTensor<Self> {
154 rudnn::convolution::tensor::conv_weight_backward::<R, 2>(
155 x,
156 output_grad,
157 weight.shape(),
158 options,
159 Default::default(),
160 )
161 .unwrap()
162 }
163
164 fn deform_conv2d(
165 x: FloatTensor<Self>,
166 offset: FloatTensor<Self>,
167 weight: FloatTensor<Self>,
168 mask: Option<FloatTensor<Self>>,
169 bias: Option<FloatTensor<Self>>,
170 options: DeformConvOptions<2>,
171 ) -> FloatTensor<Self> {
172 rudnn::convolution::tensor::deform_conv2d(x, offset, weight, mask, bias, options).unwrap()
173 }
174
175 fn deform_conv2d_backward(
176 x: FloatTensor<Self>,
177 offset: FloatTensor<Self>,
178 weight: FloatTensor<Self>,
179 mask: Option<FloatTensor<Self>>,
180 bias: Option<FloatTensor<Self>>,
181 output_grad: FloatTensor<Self>,
182 options: DeformConvOptions<2>,
183 ) -> DeformConv2dBackward<Self> {
184 let (x, o, w, m, b) = rudnn::convolution::tensor::deform_conv2d_backward(
185 x,
186 offset,
187 weight,
188 mask,
189 bias,
190 output_grad,
191 options,
192 )
193 .unwrap();
194 DeformConv2dBackward::new(x, o, w, m, b)
195 }
196
197 fn conv3d(
198 x: FloatTensor<Self>,
199 weight: FloatTensor<Self>,
200 bias: Option<FloatTensor<Self>>,
201 options: ConvOptions<3>,
202 ) -> FloatTensor<Self> {
203 rudnn::convolution::tensor::conv_forward::<R, 3>(x, weight, bias, options, Default::default()).unwrap()
204 }
205
206 fn conv3d_x_backward(
207 x: FloatTensor<Self>,
208 weight: FloatTensor<Self>,
209 output_grad: FloatTensor<Self>,
210 options: ConvOptions<3>,
211 ) -> FloatTensor<Self> {
212 rudnn::convolution::tensor::conv_data_backward(
213 output_grad,
214 weight,
215 x.shape(),
216 options,
217 Default::default(),
218 )
219 .unwrap()
220 }
221
222 fn conv3d_weight_backward(
223 x: FloatTensor<Self>,
224 weight: FloatTensor<Self>,
225 output_grad: FloatTensor<Self>,
226 options: ConvOptions<3>,
227 ) -> FloatTensor<Self> {
228 rudnn::convolution::tensor::conv_weight_backward::<R, 3>(
229 x,
230 output_grad,
231 weight.shape(),
232 options,
233 Default::default(),
234 )
235 .unwrap()
236 }
237
238 fn conv_transpose2d(
239 x: FloatTensor<Self>,
240 weight: FloatTensor<Self>,
241 bias: Option<FloatTensor<Self>>,
242 options: ConvTransposeOptions<2>,
243 ) -> FloatTensor<Self> {
244 rudnn::convolution::tensor::conv_transpose2d(x, weight, bias, options, ConvTranspose2dStrategy::default())
245 .unwrap()
246 }
247
248 fn conv_transpose3d(
249 x: FloatTensor<Self>,
250 weight: FloatTensor<Self>,
251 bias: Option<FloatTensor<Self>>,
252 options: ConvTransposeOptions<3>,
253 ) -> FloatTensor<Self> {
254 rudnn::convolution::tensor::conv_transpose3d(x, weight, bias, options).expect("Kernel to never fail")
255 }
256
257 fn avg_pool2d(
258 x: FloatTensor<Self>,
259 kernel_size: [usize; 2],
260 stride: [usize; 2],
261 padding: [usize; 2],
262 count_include_pad: bool,
263 ceil_mode: bool,
264 ) -> FloatTensor<Self> {
265 rudnn::pooling::avg_pool2d(
266 x,
267 kernel_size,
268 stride,
269 padding,
270 count_include_pad,
271 ceil_mode,
272 )
273 }
274
275 fn avg_pool2d_backward(
276 x: FloatTensor<Self>,
277 grad: FloatTensor<Self>,
278 kernel_size: [usize; 2],
279 stride: [usize; 2],
280 padding: [usize; 2],
281 count_include_pad: bool,
282 ceil_mode: bool,
283 ) -> FloatTensor<Self> {
284 rudnn::pooling::avg_pool2d_backward(
285 x,
286 grad,
287 kernel_size,
288 stride,
289 padding,
290 count_include_pad,
291 ceil_mode,
292 )
293 }
294
295 fn max_pool2d(
296 x: FloatTensor<Self>,
297 kernel_size: [usize; 2],
298 stride: [usize; 2],
299 padding: [usize; 2],
300 dilation: [usize; 2],
301 ceil_mode: bool,
302 ) -> FloatTensor<Self> {
303 rudnn::pooling::max_pool2d(x, kernel_size, stride, padding, dilation, ceil_mode)
304 }
305
306 fn max_pool2d_with_indices(
307 x: FloatTensor<Self>,
308 kernel_size: [usize; 2],
309 stride: [usize; 2],
310 padding: [usize; 2],
311 dilation: [usize; 2],
312 ceil_mode: bool,
313 ) -> MaxPool2dWithIndices<Self> {
314 let (output, indices) = rudnn::pooling::max_pool2d_with_indices(
315 x,
316 kernel_size,
317 stride,
318 padding,
319 dilation,
320 ceil_mode,
321 I::dtype(),
322 );
323
324 MaxPool2dWithIndices::new(output, indices)
325 }
326
327 fn max_pool2d_with_indices_backward(
328 x: FloatTensor<Self>,
329 kernel_size: [usize; 2],
330 stride: [usize; 2],
331 padding: [usize; 2],
332 dilation: [usize; 2],
333 ceil_mode: bool,
334 output_grad: FloatTensor<Self>,
335 indices: IntTensor<Self>,
336 ) -> MaxPool2dBackward<Self> {
337 MaxPool2dBackward::new(rudnn::pooling::max_pool2d_with_indices_backward(
338 x,
339 output_grad,
340 indices,
341 kernel_size,
342 stride,
343 padding,
344 dilation,
345 ceil_mode,
346 ))
347 }
348
349 fn adaptive_avg_pool2d(x: FloatTensor<Self>, output_size: [usize; 2]) -> FloatTensor<Self> {
350 rudnn::pooling::adaptive_avg_pool2d(x, output_size)
351 }
352
353 fn adaptive_avg_pool2d_backward(
354 x: FloatTensor<Self>,
355 grad: FloatTensor<Self>,
356 ) -> FloatTensor<Self> {
357 rudnn::pooling::adaptive_avg_pool2d_backward(x, grad)
358 }
359
360 fn interpolate(
361 x: FloatTensor<Self>,
362 output_size: [usize; 2],
363 options: InterpolateOptions,
364 ) -> FloatTensor<Self> {
365 rudnn::interpolation::interpolate(x, output_size, options)
366 }
367
368 fn interpolate_backward(
369 x: FloatTensor<Self>,
370 grad: FloatTensor<Self>,
371 output_size: [usize; 2],
372 options: InterpolateOptions,
373 ) -> FloatTensor<Self> {
374 rudnn::interpolation::interpolate_backward(x, grad, output_size, options)
375 }
376
377 fn attention(
378 query: FloatTensor<Self>,
379 key: FloatTensor<Self>,
380 value: FloatTensor<Self>,
381 mask: Option<BoolTensor<Self>>,
382 attn_bias: Option<FloatTensor<Self>>,
383 options: AttentionModuleOptions,
384 ) -> FloatTensor<Self> {
385 if attn_bias.is_some() || options.softcap.is_some() || options.scale.is_some() {
387 return ruda_tensor::ops::attention::attention_fallback::<Self>(
388 query, key, value, mask, attn_bias, options,
389 );
390 }
391
392 rudnn::attention::tensor::attention(
393 query,
394 key,
395 value,
396 mask,
397 attn_bias,
398 options,
399 Default::default(),
400 )
401 .expect("Kernel to never fail")
402 }
403
404 fn has_ctc_loss_backward() -> bool {
405 true
406 }
407
408 fn ctc_loss(
409 log_probs: FloatTensor<Self>,
410 targets: IntTensor<Self>,
411 input_lengths: IntTensor<Self>,
412 target_lengths: IntTensor<Self>,
413 blank: usize,
414 ) -> FloatTensor<Self> {
415 rudnn::ctc::ctc_loss(log_probs, targets, input_lengths, target_lengths, blank)
416 }
417
418 fn ctc_loss_backward(
419 log_probs: FloatTensor<Self>,
420 targets: IntTensor<Self>,
421 input_lengths: IntTensor<Self>,
422 target_lengths: IntTensor<Self>,
423 grad_loss: FloatTensor<Self>,
424 blank: usize,
425 ) -> FloatTensor<Self> {
426 let (log_alpha_full, log_beta_full, nll) = rudnn::ctc::ctc_alpha_beta(
427 log_probs.clone(),
428 targets.clone(),
429 input_lengths.clone(),
430 target_lengths,
431 blank,
432 );
433 ruda_tensor::ops::ctc::ctc_grad_from_alpha_beta_default::<Self>(
434 log_probs,
435 targets,
436 input_lengths,
437 grad_loss,
438 log_alpha_full,
439 log_beta_full,
440 nll,
441 blank,
442 )
443 }
444
445 fn rfft(
446 signal: FloatTensor<Self>,
447 dim: usize,
448 n: Option<usize>,
449 ) -> (FloatTensor<Self>, FloatTensor<Self>) {
450 rufft::tensor::rfft(signal, dim, n)
451 }
452
453 fn irfft(
454 spectrum_re: FloatTensor<Self>,
455 spectrum_im: FloatTensor<Self>,
456 dim: usize,
457 n: Option<usize>,
458 ) -> FloatTensor<Self> {
459 rufft::tensor::irfft(spectrum_re, spectrum_im, dim, n)
460 }
461}