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, FloatTensorOps, 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
27fn native_norm_supported<R: DeviceRuntime>(tensor: &crate::RudaTensor<R>) -> bool {
28    let properties = tensor.client.properties();
29    let hardware = &properties.hardware;
30    let plane = hardware.plane_size_max;
31    matches!(tensor.dtype, ruda_core::tensor::DType::F32 | ruda_core::tensor::DType::F16 | ruda_core::tensor::DType::BF16)
32        && tensor.meta.num_elements() <= u32::MAX as usize
33        && tensor.meta.shape().last().is_some_and(|width| *width <= u32::MAX as usize)
34        && plane.is_power_of_two() && plane == hardware.plane_size_min
35        && properties.features.plane.contains(ruda_core::ir::features::Plane::Ops)
36        && plane <= hardware.max_ruda_dim.0 && hardware.max_ruda_dim.1 >= 4
37        && plane <= hardware.max_units_per_ruda / 4
38        && tensor.meta.shape().last().is_some_and(|width| *width > 0
39            && tensor.meta.num_elements() / width <= hardware.max_ruda_count.0 as usize)
40}
41
42fn native_rms_supported<R: DeviceRuntime>(tensor: &crate::RudaTensor<R>) -> bool {
43    let properties = tensor.client.properties();
44    let hardware = &properties.hardware;
45    let plane = hardware.plane_size_max;
46    matches!(tensor.dtype, ruda_core::tensor::DType::F32 | ruda_core::tensor::DType::F16 | ruda_core::tensor::DType::BF16)
47        && tensor.meta.num_elements() <= u32::MAX as usize
48        && properties.features.plane.contains(ruda_core::ir::features::Plane::Ops)
49        && plane.is_power_of_two() && plane <= hardware.max_ruda_dim.0.min(hardware.max_units_per_ruda)
50        && tensor.meta.shape().last().is_some_and(|width| *width > 0 && *width <= u32::MAX as usize
51            && tensor.meta.num_elements() / width <= hardware.max_ruda_count.0 as usize)
52}
53
54fn native_softmax_supported<R: DeviceRuntime>(tensor: &crate::RudaTensor<R>) -> bool {
55    native_rms_supported(tensor) && tensor.qparams.is_none()
56        && tensor.client.properties().hardware.plane_size_min == tensor.client.properties().hardware.plane_size_max
57}
58
59fn group_storage_supported<R: DeviceRuntime>(tensor: &crate::RudaTensor<R>) -> bool {
60    tensor.qparams.is_none() && matches!(tensor.dtype,
61        ruda_core::tensor::DType::F32 | ruda_core::tensor::DType::F16 | ruda_core::tensor::DType::BF16)
62}
63
64fn native_group_supported<R: DeviceRuntime>(tensor: &crate::RudaTensor<R>, groups: usize) -> bool {
65    let info = ruda_tensor::ops::group_normalization::geometry(tensor.meta.shape(), groups);
66    let properties = tensor.client.properties();
67    let hardware = &properties.hardware;
68    let plane = hardware.plane_size_max;
69    group_storage_supported(tensor) && tensor.meta.num_elements() <= u32::MAX as usize
70        && info.channels <= u32::MAX as usize && info.width <= u32::MAX as usize
71        && plane.is_power_of_two() && plane == hardware.plane_size_min
72        && properties.features.plane.contains(ruda_core::ir::features::Plane::Ops)
73        && plane <= hardware.max_ruda_dim.0 && hardware.max_ruda_dim.1 >= 4
74        && plane <= hardware.max_units_per_ruda / 4 && info.rows <= hardware.max_ruda_count.0 as usize
75}
76
77impl<R, F, I, BT> ModuleOps<Self> for DeviceBackend<R, F, I, BT>
78where
79    R: DeviceRuntime,
80    F: FloatElement,
81    I: IntElement,
82    BT: BoolElement,
83{
84    fn exponential_relu_native(tensor: FloatTensor<Self>, alpha: f64, continuous: bool) -> FloatTensor<Self> {
85        if group_storage_supported(&tensor) && tensor.meta.num_elements() <= u32::MAX as usize {
86            ruprim::elementwise::unary::exponential_relu::launch(tensor, alpha as f32, continuous)
87        } else { ruda_tensor::ops::activation_training::exponential_relu_native::<Self>(tensor, alpha, continuous) }
88    }
89
90    fn exponential_relu_native_backward(tensor: FloatTensor<Self>, grad: FloatTensor<Self>, alpha: f64, continuous: bool) -> FloatTensor<Self> {
91        if [&tensor, &grad].into_iter().all(|value| group_storage_supported(value) && value.meta.num_elements() <= u32::MAX as usize) {
92            ruprim::elementwise::unary::exponential_relu::launch_backward(tensor, grad, alpha as f32, continuous)
93        } else { ruda_tensor::ops::activation_training::exponential_relu_native_backward::<Self>(tensor, grad, alpha, continuous) }
94    }
95
96    fn leaky_relu_native(tensor: FloatTensor<Self>, negative_slope: f64) -> FloatTensor<Self> {
97        if group_storage_supported(&tensor) && tensor.meta.num_elements() <= u32::MAX as usize {
98            ruprim::elementwise::unary::leaky_relu::launch(tensor, negative_slope as f32)
99        } else { ruda_tensor::ops::activation_training::leaky_relu_native::<Self>(tensor, negative_slope) }
100    }
101
102    fn leaky_relu_native_backward(tensor: FloatTensor<Self>, grad: FloatTensor<Self>, negative_slope: f64) -> FloatTensor<Self> {
103        if [&tensor, &grad].into_iter().all(|value| group_storage_supported(value) && value.meta.num_elements() <= u32::MAX as usize) {
104            ruprim::elementwise::unary::leaky_relu::launch_backward(tensor, grad, negative_slope as f32)
105        } else { ruda_tensor::ops::activation_training::leaky_relu_native_backward::<Self>(tensor, grad, negative_slope) }
106    }
107
108    fn prelu_native(tensor: FloatTensor<Self>, alpha: FloatTensor<Self>) -> FloatTensor<Self> {
109        let info = ruda_tensor::ops::prelu_training::geometry(tensor.meta.shape(), alpha.meta.shape());
110        if group_storage_supported(&tensor) && group_storage_supported(&alpha) && info.elements <= u32::MAX as usize
111            && info.parameters <= u32::MAX as usize && info.channels <= u32::MAX as usize && info.spatial <= u32::MAX as usize {
112            ruprim::elementwise::unary::prelu::launch(tensor, alpha)
113        } else { ruda_tensor::ops::prelu_training::prelu_native::<Self>(tensor, alpha) }
114    }
115
116    fn prelu_native_backward_select(tensor: FloatTensor<Self>, alpha: FloatTensor<Self>, grad: FloatTensor<Self>,
117        mask: [bool; 2]) -> [Option<FloatTensor<Self>>; 2] {
118        if mask == [false; 2] { return [None, None]; }
119        let info = ruda_tensor::ops::prelu_training::geometry(tensor.meta.shape(), alpha.meta.shape());
120        if [&tensor, &alpha, &grad].into_iter().all(group_storage_supported) && info.elements <= u32::MAX as usize
121            && info.parameters <= u32::MAX as usize && info.channels <= u32::MAX as usize && info.spatial <= u32::MAX as usize {
122            ruprim::elementwise::unary::prelu::launch_backward_select(tensor, alpha, grad, mask)
123        } else { ruda_tensor::ops::prelu_training::prelu_native_backward_select::<Self>(tensor, alpha, grad, mask) }
124    }
125
126    fn group_norm_with_stats(tensor: FloatTensor<Self>, gamma: Option<FloatTensor<Self>>,
127        beta: Option<FloatTensor<Self>>, groups: usize, epsilon: f64) -> ruda_tensor::ops::LayerNormOutput<Self> {
128        if native_group_supported(&tensor, groups) && gamma.iter().chain(beta.iter()).all(group_storage_supported) {
129            let [output, mean, rstd] = rudnn::normalization::group_norm_with_stats(tensor, gamma, beta, groups, epsilon as f32)
130                .expect("invalid native GroupNorm bindings");
131            ruda_tensor::ops::LayerNormOutput { output, mean, rstd }
132        } else { ruda_tensor::ops::group_normalization::group_norm_with_stats::<Self>(tensor, gamma, beta, groups, epsilon) }
133    }
134
135    fn group_norm_backward_select(tensor: FloatTensor<Self>, gamma: Option<FloatTensor<Self>>, grad: FloatTensor<Self>,
136        mean: FloatTensor<Self>, rstd: FloatTensor<Self>, groups: usize, mask: [bool; 3]) -> [Option<FloatTensor<Self>>; 3] {
137        if mask == [false; 3] { return [None, None, None]; }
138        if native_group_supported(&tensor, groups) && group_storage_supported(&grad) && gamma.iter().all(group_storage_supported)
139            && mean.dtype == ruda_core::tensor::DType::F32 && rstd.dtype == ruda_core::tensor::DType::F32 {
140            rudnn::normalization::group_norm_backward_select(tensor, gamma, grad, mean, rstd, groups, mask)
141                .expect("invalid native GroupNorm backward bindings")
142        } else { ruda_tensor::ops::group_normalization::group_norm_backward_select::<Self>(tensor, gamma, grad, mean, rstd, groups, mask) }
143    }
144
145    fn gelu_native(tensor: FloatTensor<Self>, approximate: bool) -> FloatTensor<Self> {
146        if matches!(tensor.dtype, ruda_core::tensor::DType::F32 | ruda_core::tensor::DType::F16 | ruda_core::tensor::DType::BF16)
147            && tensor.qparams.is_none() {
148            ruprim::elementwise::unary::gelu::launch(tensor, approximate)
149        } else { ruda_tensor::ops::activation_training::gelu_native::<Self>(tensor, approximate) }
150    }
151
152    fn gelu_native_backward(input: FloatTensor<Self>, grad: FloatTensor<Self>, approximate: bool) -> FloatTensor<Self> {
153        if [&input, &grad].iter().all(|value| value.qparams.is_none()
154            && matches!(value.dtype, ruda_core::tensor::DType::F32 | ruda_core::tensor::DType::F16 | ruda_core::tensor::DType::BF16)) {
155            ruprim::elementwise::unary::gelu::launch_backward(input, grad, approximate)
156        } else { ruda_tensor::ops::activation_training::gelu_native_backward::<Self>(input, grad, approximate) }
157    }
158
159    fn silu_native(tensor: FloatTensor<Self>) -> FloatTensor<Self> {
160        if matches!(tensor.dtype, ruda_core::tensor::DType::F32 | ruda_core::tensor::DType::F16 | ruda_core::tensor::DType::BF16)
161            && tensor.qparams.is_none() {
162            ruprim::elementwise::unary::silu::launch(tensor)
163        } else { ruda_tensor::ops::activation_training::silu_native::<Self>(tensor) }
164    }
165
166    fn silu_native_backward(input: FloatTensor<Self>, grad: FloatTensor<Self>) -> FloatTensor<Self> {
167        if [&input, &grad].iter().all(|value| value.qparams.is_none()
168            && matches!(value.dtype, ruda_core::tensor::DType::F32 | ruda_core::tensor::DType::F16 | ruda_core::tensor::DType::BF16)) {
169            ruprim::elementwise::unary::silu::launch_backward(input, grad)
170        } else { ruda_tensor::ops::activation_training::silu_native_backward::<Self>(input, grad) }
171    }
172
173    fn softmax_with_stats(tensor: FloatTensor<Self>, dim: usize, logarithmic: bool)
174        -> ruda_tensor::ops::SoftmaxOutput<Self> {
175        let rank = tensor.meta.shape().num_dims();
176        assert!(dim < rank, "softmax axis out of bounds");
177        assert!(tensor.meta.shape()[dim] > 0, "softmax axis must be nonempty");
178        let storage: ruda_core::tensor::FloatDType = tensor.dtype.into();
179        let tensor = if dim == rank - 1 { tensor } else { Self::float_swap_dims(tensor, dim, rank - 1) };
180        let result = if native_softmax_supported(&tensor) {
181            let working = rudnn::normalization::softmax_last_axis_working(tensor, logarithmic)
182                .expect("invalid native softmax bindings");
183            ruda_tensor::ops::SoftmaxOutput { output: Self::float_cast(working.clone(), storage), working }
184        } else { ruda_tensor::ops::softmax::softmax_with_stats::<Self>(tensor, rank - 1, logarithmic) };
185        if dim == rank - 1 { result } else {
186            ruda_tensor::ops::SoftmaxOutput {
187                output: Self::float_swap_dims(result.output, dim, rank - 1),
188                working: Self::float_swap_dims(result.working, dim, rank - 1),
189            }
190        }
191    }
192
193    fn softmax_native_backward(working: FloatTensor<Self>, grad: FloatTensor<Self>, dim: usize,
194        logarithmic: bool) -> FloatTensor<Self> {
195        let rank = working.meta.shape().num_dims();
196        assert!(dim < rank, "softmax backward axis out of bounds");
197        assert_eq!(grad.meta.shape(), working.meta.shape(), "softmax gradient shape differs");
198        let working = if dim == rank - 1 { working } else { Self::float_swap_dims(working, dim, rank - 1) };
199        let grad = if dim == rank - 1 { grad } else { Self::float_swap_dims(grad, dim, rank - 1) };
200        let result = if working.dtype == ruda_core::tensor::DType::F32
201            && native_softmax_supported(&working) && native_softmax_supported(&grad) {
202            rudnn::normalization::softmax_last_axis_backward(working, grad, logarithmic)
203                .expect("invalid native softmax backward bindings")
204        } else { ruda_tensor::ops::softmax::softmax_backward::<Self>(working, grad, rank - 1, logarithmic) };
205        if dim == rank - 1 { result } else { Self::float_swap_dims(result, dim, rank - 1) }
206    }
207
208    fn has_layer_norm_backward() -> bool { true }
209
210    fn layer_norm_backward_select(tensor: FloatTensor<Self>, gamma: FloatTensor<Self>, grad: FloatTensor<Self>,
211        mean: FloatTensor<Self>, rstd: FloatTensor<Self>, mask: [bool; 3]) -> [Option<FloatTensor<Self>>; 3] {
212        if mask == [false; 3] { return [None, None, None]; }
213        if R::has_native_layer_norm() {
214            let out = Self::layer_norm_backward(tensor, gamma, grad, mean, rstd);
215            return core::array::from_fn(|index| if mask[index] {
216                Some(match index { 0 => out.input.clone(), 1 => out.weight.clone(), _ => out.bias.clone() })
217            } else { None });
218        }
219        if native_norm_supported(&tensor) && native_norm_supported(&gamma) && native_norm_supported(&grad)
220            && mean.dtype == ruda_core::tensor::DType::F32 && rstd.dtype == ruda_core::tensor::DType::F32 {
221            return rudnn::normalization::layer_norm_backward_select(tensor, gamma, grad, mean, rstd, mask)
222                .expect("invalid native LayerNorm backward bindings");
223        }
224        ruda_tensor::ops::normalization::layer_norm_backward_select::<Self>(tensor, gamma, grad, mean, rstd, mask)
225    }
226
227    fn has_rms_norm_backward() -> bool { true }
228
229    fn rms_norm_backward_select(tensor: FloatTensor<Self>, gamma: FloatTensor<Self>, grad: FloatTensor<Self>,
230        rstd: FloatTensor<Self>, mask: [bool; 2]) -> [Option<FloatTensor<Self>>; 2] {
231        if mask == [false; 2] { return [None, None]; }
232        if native_rms_supported(&tensor) && native_rms_supported(&gamma) && native_rms_supported(&grad)
233            && rstd.dtype == ruda_core::tensor::DType::F32 {
234            return rudnn::normalization::rms_norm_backward_select(tensor, gamma, grad, rstd, mask)
235                .expect("invalid native RMSNorm backward bindings");
236        }
237        ruda_tensor::ops::normalization::rms_norm_backward_select::<Self>(tensor, gamma, grad, rstd, mask)
238    }
239
240    fn rms_norm_with_stats(tensor: FloatTensor<Self>, gamma: FloatTensor<Self>, epsilon: f64)
241        -> ruda_tensor::ops::RmsNormOutput<Self> {
242        if native_rms_supported(&tensor) && native_rms_supported(&gamma)
243            && (epsilon as f32).is_finite() && (epsilon as f32) > 0.0 {
244            let [output, rstd] = rudnn::normalization::rms_norm_with_stats(tensor, gamma, epsilon as f32)
245                .expect("invalid native RMSNorm bindings");
246            return ruda_tensor::ops::RmsNormOutput { output, rstd };
247        }
248        ruda_tensor::ops::normalization::rms_norm_with_stats::<Self>(tensor, gamma, epsilon)
249    }
250
251    fn rms_norm_backward(tensor: FloatTensor<Self>, gamma: FloatTensor<Self>, grad: FloatTensor<Self>,
252        rstd: FloatTensor<Self>) -> ruda_tensor::ops::RmsNormBackward<Self> {
253        if native_rms_supported(&tensor) && native_rms_supported(&gamma) && native_rms_supported(&grad)
254            && rstd.dtype == ruda_core::tensor::DType::F32 {
255            let [input, weight] = rudnn::normalization::rms_norm_backward(tensor, gamma, grad, rstd)
256                .expect("invalid native RMSNorm backward bindings");
257            return ruda_tensor::ops::RmsNormBackward { input, weight };
258        }
259        ruda_tensor::ops::normalization::rms_norm_backward::<Self>(tensor, gamma, grad, rstd)
260    }
261
262    fn layer_norm(
263        tensor: FloatTensor<Self>, gamma: FloatTensor<Self>,
264        beta: Option<FloatTensor<Self>>, epsilon: f64,
265    ) -> FloatTensor<Self> {
266        Self::layer_norm_with_stats(tensor, gamma, beta, epsilon).output
267    }
268
269    fn layer_norm_with_stats(
270        tensor: FloatTensor<Self>, gamma: FloatTensor<Self>,
271        beta: Option<FloatTensor<Self>>, epsilon: f64,
272    ) -> ruda_tensor::ops::LayerNormOutput<Self> {
273        if !R::has_native_layer_norm() {
274            if native_norm_supported(&tensor) && native_norm_supported(&gamma)
275                && beta.as_ref().is_none_or(native_norm_supported)
276                && (epsilon as f32).is_finite() && (epsilon as f32) > 0.0 {
277                let [output, mean, rstd] = rudnn::normalization::layer_norm_with_stats(
278                    tensor, gamma, beta, epsilon as f32,
279                ).expect("invalid native LayerNorm bindings");
280                return ruda_tensor::ops::LayerNormOutput { output, mean, rstd };
281            }
282            return ruda_tensor::ops::normalization::layer_norm_with_stats::<Self>(tensor, gamma, beta, epsilon);
283        }
284        let client = tensor.client.clone();
285        let device = tensor.device.clone();
286        for other in core::iter::once(&gamma).chain(beta.iter()) {
287            assert_eq!(device, other.device, "LayerNorm device mismatch");
288            assert!(client.same_execution_queue(&other.client), "LayerNorm queue mismatch");
289        }
290        let [output, mean, rstd] = R::layer_norm(
291            &client, norm_buffer(tensor), norm_buffer(gamma), beta.map(norm_buffer), epsilon,
292        );
293        let from = |buffer| norm_tensor(buffer, client.clone(), device.clone());
294        ruda_tensor::ops::LayerNormOutput { output: from(output), mean: from(mean), rstd: from(rstd) }
295    }
296
297    fn layer_norm_backward(
298        tensor: FloatTensor<Self>, gamma: FloatTensor<Self>, grad: FloatTensor<Self>,
299        mean: FloatTensor<Self>, rstd: FloatTensor<Self>,
300    ) -> ruda_tensor::ops::LayerNormBackward<Self> {
301        if !R::has_native_layer_norm() {
302            if native_norm_supported(&tensor) && native_norm_supported(&gamma) && native_norm_supported(&grad)
303                && mean.dtype == ruda_core::tensor::DType::F32 && rstd.dtype == ruda_core::tensor::DType::F32 {
304                let [input, weight, bias] = rudnn::normalization::layer_norm_backward(
305                    tensor, gamma, grad, mean, rstd,
306                ).expect("invalid native LayerNorm backward bindings");
307                return ruda_tensor::ops::LayerNormBackward { input, weight, bias };
308            }
309            return ruda_tensor::ops::normalization::layer_norm_backward::<Self>(tensor, gamma, grad, mean, rstd);
310        }
311        let client = tensor.client.clone();
312        let device = tensor.device.clone();
313        for other in [&gamma, &grad, &mean, &rstd] {
314            assert_eq!(device, other.device, "LayerNorm backward device mismatch");
315            assert!(client.same_execution_queue(&other.client), "LayerNorm backward queue mismatch");
316        }
317        let [input, weight, bias] = R::layer_norm_backward(
318            &client, norm_buffer(tensor), norm_buffer(gamma), norm_buffer(grad),
319            norm_buffer(mean), norm_buffer(rstd),
320        );
321        let from = |buffer| norm_tensor(buffer, client.clone(), device.clone());
322        ruda_tensor::ops::LayerNormBackward { input: from(input), weight: from(weight), bias: from(bias) }
323    }
324
325    fn conv1d(
326        x: FloatTensor<Self>,
327        weight: FloatTensor<Self>,
328        bias: Option<FloatTensor<Self>>,
329        options: ConvOptions<1>,
330    ) -> FloatTensor<Self> {
331        rudnn::convolution::tensor::conv_forward::<R, 1>(x, weight, bias, options, Default::default()).unwrap()
332    }
333
334    fn conv1d_x_backward(
335        x: FloatTensor<Self>,
336        weight: FloatTensor<Self>,
337        output_grad: FloatTensor<Self>,
338        options: ConvOptions<1>,
339    ) -> FloatTensor<Self> {
340        rudnn::convolution::tensor::conv_data_backward(
341            output_grad,
342            weight,
343            x.shape(),
344            options,
345            Default::default(),
346        )
347        .unwrap()
348    }
349
350    fn conv1d_weight_backward(
351        x: FloatTensor<Self>,
352        weight: FloatTensor<Self>,
353        output_grad: FloatTensor<Self>,
354        options: ConvOptions<1>,
355    ) -> FloatTensor<Self> {
356        rudnn::convolution::tensor::conv_weight_backward::<R, 1>(
357            x,
358            output_grad,
359            weight.shape(),
360            options,
361            Default::default(),
362        )
363        .unwrap()
364    }
365
366    fn conv2d(
367        x: FloatTensor<Self>,
368        weight: FloatTensor<Self>,
369        bias: Option<FloatTensor<Self>>,
370        options: ConvOptions<2>,
371    ) -> FloatTensor<Self> {
372        rudnn::convolution::tensor::conv_forward::<R, 2>(x, weight, bias, options, Default::default()).unwrap()
373    }
374
375    fn conv2d_x_backward(
376        x: FloatTensor<Self>,
377        weight: FloatTensor<Self>,
378        output_grad: FloatTensor<Self>,
379        options: ConvOptions<2>,
380    ) -> FloatTensor<Self> {
381        rudnn::convolution::tensor::conv_data_backward(
382            output_grad,
383            weight,
384            x.shape(),
385            options,
386            Default::default(),
387        )
388        .unwrap()
389    }
390
391    fn conv2d_weight_backward(
392        x: FloatTensor<Self>,
393        weight: FloatTensor<Self>,
394        output_grad: FloatTensor<Self>,
395        options: ConvOptions<2>,
396    ) -> FloatTensor<Self> {
397        rudnn::convolution::tensor::conv_weight_backward::<R, 2>(
398            x,
399            output_grad,
400            weight.shape(),
401            options,
402            Default::default(),
403        )
404        .unwrap()
405    }
406
407    fn deform_conv2d(
408        x: FloatTensor<Self>,
409        offset: FloatTensor<Self>,
410        weight: FloatTensor<Self>,
411        mask: Option<FloatTensor<Self>>,
412        bias: Option<FloatTensor<Self>>,
413        options: DeformConvOptions<2>,
414    ) -> FloatTensor<Self> {
415        rudnn::convolution::tensor::deform_conv2d(x, offset, weight, mask, bias, options).unwrap()
416    }
417
418    fn deform_conv2d_backward(
419        x: FloatTensor<Self>,
420        offset: FloatTensor<Self>,
421        weight: FloatTensor<Self>,
422        mask: Option<FloatTensor<Self>>,
423        bias: Option<FloatTensor<Self>>,
424        output_grad: FloatTensor<Self>,
425        options: DeformConvOptions<2>,
426    ) -> DeformConv2dBackward<Self> {
427        let (x, o, w, m, b) = rudnn::convolution::tensor::deform_conv2d_backward(
428            x,
429            offset,
430            weight,
431            mask,
432            bias,
433            output_grad,
434            options,
435        )
436        .unwrap();
437        DeformConv2dBackward::new(x, o, w, m, b)
438    }
439
440    fn conv3d(
441        x: FloatTensor<Self>,
442        weight: FloatTensor<Self>,
443        bias: Option<FloatTensor<Self>>,
444        options: ConvOptions<3>,
445    ) -> FloatTensor<Self> {
446        rudnn::convolution::tensor::conv_forward::<R, 3>(x, weight, bias, options, Default::default()).unwrap()
447    }
448
449    fn conv3d_x_backward(
450        x: FloatTensor<Self>,
451        weight: FloatTensor<Self>,
452        output_grad: FloatTensor<Self>,
453        options: ConvOptions<3>,
454    ) -> FloatTensor<Self> {
455        rudnn::convolution::tensor::conv_data_backward(
456            output_grad,
457            weight,
458            x.shape(),
459            options,
460            Default::default(),
461        )
462        .unwrap()
463    }
464
465    fn conv3d_weight_backward(
466        x: FloatTensor<Self>,
467        weight: FloatTensor<Self>,
468        output_grad: FloatTensor<Self>,
469        options: ConvOptions<3>,
470    ) -> FloatTensor<Self> {
471        rudnn::convolution::tensor::conv_weight_backward::<R, 3>(
472            x,
473            output_grad,
474            weight.shape(),
475            options,
476            Default::default(),
477        )
478        .unwrap()
479    }
480
481    fn conv_transpose2d(
482        x: FloatTensor<Self>,
483        weight: FloatTensor<Self>,
484        bias: Option<FloatTensor<Self>>,
485        options: ConvTransposeOptions<2>,
486    ) -> FloatTensor<Self> {
487        rudnn::convolution::tensor::conv_transpose2d(x, weight, bias, options, ConvTranspose2dStrategy::default())
488            .unwrap()
489    }
490
491    fn conv_transpose3d(
492        x: FloatTensor<Self>,
493        weight: FloatTensor<Self>,
494        bias: Option<FloatTensor<Self>>,
495        options: ConvTransposeOptions<3>,
496    ) -> FloatTensor<Self> {
497        rudnn::convolution::tensor::conv_transpose3d(x, weight, bias, options).expect("Kernel to never fail")
498    }
499
500    fn avg_pool2d(
501        x: FloatTensor<Self>,
502        kernel_size: [usize; 2],
503        stride: [usize; 2],
504        padding: [usize; 2],
505        count_include_pad: bool,
506        ceil_mode: bool,
507    ) -> FloatTensor<Self> {
508        rudnn::pooling::avg_pool2d(
509            x,
510            kernel_size,
511            stride,
512            padding,
513            count_include_pad,
514            ceil_mode,
515        )
516    }
517
518    fn avg_pool2d_backward(
519        x: FloatTensor<Self>,
520        grad: FloatTensor<Self>,
521        kernel_size: [usize; 2],
522        stride: [usize; 2],
523        padding: [usize; 2],
524        count_include_pad: bool,
525        ceil_mode: bool,
526    ) -> FloatTensor<Self> {
527        rudnn::pooling::avg_pool2d_backward(
528            x,
529            grad,
530            kernel_size,
531            stride,
532            padding,
533            count_include_pad,
534            ceil_mode,
535        )
536    }
537
538    fn max_pool2d(
539        x: FloatTensor<Self>,
540        kernel_size: [usize; 2],
541        stride: [usize; 2],
542        padding: [usize; 2],
543        dilation: [usize; 2],
544        ceil_mode: bool,
545    ) -> FloatTensor<Self> {
546        rudnn::pooling::max_pool2d(x, kernel_size, stride, padding, dilation, ceil_mode)
547    }
548
549    fn max_pool2d_with_indices(
550        x: FloatTensor<Self>,
551        kernel_size: [usize; 2],
552        stride: [usize; 2],
553        padding: [usize; 2],
554        dilation: [usize; 2],
555        ceil_mode: bool,
556    ) -> MaxPool2dWithIndices<Self> {
557        let (output, indices) = rudnn::pooling::max_pool2d_with_indices(
558            x,
559            kernel_size,
560            stride,
561            padding,
562            dilation,
563            ceil_mode,
564            I::dtype(),
565        );
566
567        MaxPool2dWithIndices::new(output, indices)
568    }
569
570    fn max_pool2d_with_indices_backward(
571        x: FloatTensor<Self>,
572        kernel_size: [usize; 2],
573        stride: [usize; 2],
574        padding: [usize; 2],
575        dilation: [usize; 2],
576        ceil_mode: bool,
577        output_grad: FloatTensor<Self>,
578        indices: IntTensor<Self>,
579    ) -> MaxPool2dBackward<Self> {
580        MaxPool2dBackward::new(rudnn::pooling::max_pool2d_with_indices_backward(
581            x,
582            output_grad,
583            indices,
584            kernel_size,
585            stride,
586            padding,
587            dilation,
588            ceil_mode,
589        ))
590    }
591
592    fn adaptive_avg_pool2d(x: FloatTensor<Self>, output_size: [usize; 2]) -> FloatTensor<Self> {
593        rudnn::pooling::adaptive_avg_pool2d(x, output_size)
594    }
595
596    fn adaptive_avg_pool3d(x: FloatTensor<Self>, output_size: [usize; 3]) -> FloatTensor<Self> {
597        rudnn::pooling::adaptive_avg_pool3d(x, output_size)
598    }
599
600    fn max_pool3d(x: FloatTensor<Self>, kernel: [usize; 3], stride: [usize; 3],
601        padding: [usize; 3], dilation: [usize; 3], ceil: bool) -> FloatTensor<Self> {
602        rudnn::pooling::max_pool3d(x, kernel, stride, padding, dilation, ceil)
603    }
604
605    fn max_pool3d_with_indices(x: FloatTensor<Self>, kernel: [usize; 3], stride: [usize; 3],
606        padding: [usize; 3], dilation: [usize; 3], ceil: bool) -> ruda_tensor::ops::MaxPool3dWithIndices<Self> {
607        let (output, indices) = rudnn::pooling::max_pool3d_with_indices(x, kernel, stride, padding, dilation, ceil);
608        ruda_tensor::ops::MaxPool3dWithIndices::new(output, indices)
609    }
610
611    fn max_pool3d_with_indices_backward(x: FloatTensor<Self>, grad: FloatTensor<Self>, indices: IntTensor<Self>,
612        kernel: [usize; 3], stride: [usize; 3], padding: [usize; 3], dilation: [usize; 3],
613        ceil: bool) -> ruda_tensor::ops::MaxPool3dBackward<Self> {
614        ruda_tensor::ops::MaxPool3dBackward::new(rudnn::pooling::max_pool3d_with_indices_backward(
615            x, grad, indices, kernel, stride, padding, dilation, ceil))
616    }
617
618    fn avg_pool3d_native_output_size(input: [usize; 3], kernel: [usize; 3],
619        stride: [usize; 3], padding: [usize; 3], ceil: bool) -> Option<[usize; 3]> {
620        Some(rudnn::pooling::avg_pool3d_output_size(input, kernel, stride, padding, ceil))
621    }
622
623    fn avg_pool3d(x: FloatTensor<Self>, kernel: [usize; 3], stride: [usize; 3],
624        padding: [usize; 3], include_pad: bool, ceil: bool) -> FloatTensor<Self> {
625        rudnn::pooling::avg_pool3d(x, kernel, stride, padding, include_pad, ceil)
626    }
627
628    fn avg_pool3d_backward(x: FloatTensor<Self>, grad: FloatTensor<Self>, kernel: [usize; 3],
629        stride: [usize; 3], padding: [usize; 3], include_pad: bool, ceil: bool) -> FloatTensor<Self> {
630        rudnn::pooling::avg_pool3d_backward(x, grad, kernel, stride, padding, include_pad, ceil)
631    }
632
633    fn adaptive_avg_pool3d_backward(x: FloatTensor<Self>, grad: FloatTensor<Self>) -> FloatTensor<Self> {
634        rudnn::pooling::adaptive_avg_pool3d_backward(x, grad)
635    }
636
637    fn adaptive_avg_pool2d_backward(
638        x: FloatTensor<Self>,
639        grad: FloatTensor<Self>,
640    ) -> FloatTensor<Self> {
641        rudnn::pooling::adaptive_avg_pool2d_backward(x, grad)
642    }
643
644    fn interpolate(
645        x: FloatTensor<Self>,
646        output_size: [usize; 2],
647        options: InterpolateOptions,
648    ) -> FloatTensor<Self> {
649        rudnn::interpolation::interpolate(x, output_size, options)
650    }
651
652    fn interpolate_backward(
653        x: FloatTensor<Self>,
654        grad: FloatTensor<Self>,
655        output_size: [usize; 2],
656        options: InterpolateOptions,
657    ) -> FloatTensor<Self> {
658        rudnn::interpolation::interpolate_backward(x, grad, output_size, options)
659    }
660
661    fn interpolate1d(x: FloatTensor<Self>, size: usize, options: InterpolateOptions) -> FloatTensor<Self> {
662        rudnn::interpolation::interpolate1d(x, size, options)
663    }
664
665    fn interpolate1d_backward(x: FloatTensor<Self>, grad: FloatTensor<Self>, size: usize,
666        options: InterpolateOptions) -> FloatTensor<Self> {
667        rudnn::interpolation::interpolate1d_backward(x, grad, size, options)
668    }
669
670    fn interpolate3d(x: FloatTensor<Self>, size: [usize; 3], options: InterpolateOptions) -> FloatTensor<Self> {
671        rudnn::interpolation::interpolate3d(x, size, options)
672    }
673
674    fn interpolate3d_backward(x: FloatTensor<Self>, grad: FloatTensor<Self>, size: [usize; 3],
675        options: InterpolateOptions) -> FloatTensor<Self> {
676        rudnn::interpolation::interpolate3d_backward(x, grad, size, options)
677    }
678
679    fn attention(
680        query: FloatTensor<Self>,
681        key: FloatTensor<Self>,
682        value: FloatTensor<Self>,
683        mask: Option<BoolTensor<Self>>,
684        attn_bias: Option<FloatTensor<Self>>,
685        options: AttentionModuleOptions,
686    ) -> FloatTensor<Self> {
687        // Fall back to naive attention for features the flash kernel doesn't support.
688        if attn_bias.is_some() || options.softcap.is_some() || options.scale.is_some() {
689            return ruda_tensor::ops::attention::attention_fallback::<Self>(
690                query, key, value, mask, attn_bias, options,
691            );
692        }
693
694        rudnn::attention::tensor::attention(
695            query,
696            key,
697            value,
698            mask,
699            attn_bias,
700            options,
701            Default::default(),
702        )
703        .expect("Kernel to never fail")
704    }
705
706    fn has_ctc_loss_backward() -> bool {
707        true
708    }
709
710    fn ctc_loss(
711        log_probs: FloatTensor<Self>,
712        targets: IntTensor<Self>,
713        input_lengths: IntTensor<Self>,
714        target_lengths: IntTensor<Self>,
715        blank: usize,
716    ) -> FloatTensor<Self> {
717        rudnn::ctc::ctc_loss(log_probs, targets, input_lengths, target_lengths, blank)
718    }
719
720    fn ctc_loss_backward(
721        log_probs: FloatTensor<Self>,
722        targets: IntTensor<Self>,
723        input_lengths: IntTensor<Self>,
724        target_lengths: IntTensor<Self>,
725        grad_loss: FloatTensor<Self>,
726        blank: usize,
727    ) -> FloatTensor<Self> {
728        let (log_alpha_full, log_beta_full, nll) = rudnn::ctc::ctc_alpha_beta(
729            log_probs.clone(),
730            targets.clone(),
731            input_lengths.clone(),
732            target_lengths,
733            blank,
734        );
735        ruda_tensor::ops::ctc::ctc_grad_from_alpha_beta_default::<Self>(
736            log_probs,
737            targets,
738            input_lengths,
739            grad_loss,
740            log_alpha_full,
741            log_beta_full,
742            nll,
743            blank,
744        )
745    }
746
747    fn rfft(
748        signal: FloatTensor<Self>,
749        dim: usize,
750        n: Option<usize>,
751    ) -> (FloatTensor<Self>, FloatTensor<Self>) {
752        rufft::tensor::rfft(signal, dim, n)
753    }
754
755    fn irfft(
756        spectrum_re: FloatTensor<Self>,
757        spectrum_im: FloatTensor<Self>,
758        dim: usize,
759        n: Option<usize>,
760    ) -> FloatTensor<Self> {
761        rufft::tensor::irfft(spectrum_re, spectrum_im, dim, n)
762    }
763}