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