Skip to main content

ruda_tensor_device/dispatch/
neural.rs

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        // Fall back to naive attention for features the flash kernel doesn't support.
386        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}