1use crate::api::{
2 Bool, Int, Tensor, TensorPrimitive,
3 backend::Backend,
4 check,
5 check::TensorCheck,
6 ops::{
7 AttentionModuleOptions, ConvOptions, ConvTransposeOptions, InterpolateOptions, PadMode,
8 PaddedConvOptions, UnfoldOptions,
9 },
10};
11
12use super::ops::DeformConvOptions;
13pub use super::spatial_pool::{adaptive_avg_pool3d, avg_pool3d, max_pool3d, max_pool3d_with_indices};
14pub use super::spatial_interpolate::{interpolate1d, interpolate3d};
15pub use super::spatial_pool::{
16 avg_pool1d_padded, avg_pool2d_padded, avg_pool3d_padded,
17 max_pool1d_padded, max_pool2d_padded, max_pool3d_padded,
18 max_pool1d_with_indices_padded, max_pool2d_with_indices_padded,
19 max_pool3d_with_indices_padded,
20};
21
22pub fn ctc_loss<B>(
36 log_probs: Tensor<B, 3>,
37 targets: Tensor<B, 2, Int>,
38 input_lengths: Tensor<B, 1, Int>,
39 target_lengths: Tensor<B, 1, Int>,
40 blank: usize,
41) -> Tensor<B, 1>
42where
43 B: Backend,
44{
45 Tensor::new(TensorPrimitive::Float(B::ctc_loss(
46 log_probs.primitive.tensor(),
47 targets.primitive,
48 input_lengths.primitive,
49 target_lengths.primitive,
50 blank,
51 )))
52}
53
54pub fn embedding<B>(weights: Tensor<B, 2>, indices: Tensor<B, 2, Int>) -> Tensor<B, 3>
56where
57 B: Backend,
58{
59 Tensor::new(TensorPrimitive::Float(B::embedding(
60 weights.primitive.tensor(),
61 indices.primitive,
62 )))
63}
64
65pub fn conv1d<B>(
71 x: Tensor<B, 3>,
72 weight: Tensor<B, 3>,
73 bias: Option<Tensor<B, 1>>,
74 options: impl Into<PaddedConvOptions<1>>,
75) -> Tensor<B, 3>
76where
77 B: Backend,
78{
79 let padded_options = options.into();
80 check!(TensorCheck::conv(
81 "conv1d",
82 x.dims(),
83 weight.dims(),
84 padded_options.options.groups,
85 ));
86
87 if let Some(padding_end) = padded_options.padding_end {
88 let left = padded_options.options.padding[0];
89 let right = padding_end[0];
90 let padded = x.pad((left, right, 0, 0), PadMode::Constant(0.0));
92 let zero_options = ConvOptions::new(
93 padded_options.options.stride,
94 [0],
95 padded_options.options.dilation,
96 padded_options.options.groups,
97 );
98 Tensor::new(TensorPrimitive::Float(B::conv1d(
99 padded.primitive.tensor(),
100 weight.primitive.tensor(),
101 bias.map(|b| b.primitive.tensor()),
102 zero_options,
103 )))
104 } else {
105 Tensor::new(TensorPrimitive::Float(B::conv1d(
106 x.primitive.tensor(),
107 weight.primitive.tensor(),
108 bias.map(|b| b.primitive.tensor()),
109 padded_options.options,
110 )))
111 }
112}
113
114pub fn conv2d<B>(
120 x: Tensor<B, 4>,
121 weight: Tensor<B, 4>,
122 bias: Option<Tensor<B, 1>>,
123 options: impl Into<PaddedConvOptions<2>>,
124) -> Tensor<B, 4>
125where
126 B: Backend,
127{
128 let padded_options = options.into();
129 check!(TensorCheck::conv(
130 "conv2d",
131 x.dims(),
132 weight.dims(),
133 padded_options.options.groups,
134 ));
135
136 if let Some(padding_end) = padded_options.padding_end {
137 let top = padded_options.options.padding[0];
138 let left = padded_options.options.padding[1];
139 let bottom = padding_end[0];
140 let right = padding_end[1];
141 let padded = x.pad((left, right, top, bottom), PadMode::Constant(0.0));
143 let zero_options = ConvOptions::new(
144 padded_options.options.stride,
145 [0, 0],
146 padded_options.options.dilation,
147 padded_options.options.groups,
148 );
149 Tensor::new(TensorPrimitive::Float(B::conv2d(
150 padded.primitive.tensor(),
151 weight.primitive.tensor(),
152 bias.map(|b| b.primitive.tensor()),
153 zero_options,
154 )))
155 } else {
156 Tensor::new(TensorPrimitive::Float(B::conv2d(
157 x.primitive.tensor(),
158 weight.primitive.tensor(),
159 bias.map(|b| b.primitive.tensor()),
160 padded_options.options,
161 )))
162 }
163}
164
165pub fn conv3d<B>(
171 x: Tensor<B, 5>,
172 weight: Tensor<B, 5>,
173 bias: Option<Tensor<B, 1>>,
174 options: impl Into<PaddedConvOptions<3>>,
175) -> Tensor<B, 5>
176where
177 B: Backend,
178{
179 let padded_options = options.into();
180 check!(TensorCheck::conv(
181 "conv3d",
182 x.dims(),
183 weight.dims(),
184 padded_options.options.groups,
185 ));
186
187 let mut options = padded_options.options;
188 let x = if let Some(padding_end) = padded_options.padding_end {
189 let padding: [(usize, usize); 3] =
190 core::array::from_fn(|axis| (options.padding[axis], padding_end[axis]));
191 options.padding = [0; 3];
192 x.pad(padding, PadMode::Constant(0.0))
193 } else {
194 x
195 };
196
197 Tensor::new(TensorPrimitive::Float(B::conv3d(
198 x.primitive.tensor(),
199 weight.primitive.tensor(),
200 bias.map(|b| b.primitive.tensor()),
201 options,
202 )))
203}
204
205pub fn deform_conv2d<B>(
207 x: Tensor<B, 4>,
208 offset: Tensor<B, 4>,
209 weight: Tensor<B, 4>,
210 mask: Option<Tensor<B, 4>>,
211 bias: Option<Tensor<B, 1>>,
212 options: DeformConvOptions<2>,
213) -> Tensor<B, 4>
214where
215 B: Backend,
216{
217 check!(TensorCheck::conv(
218 "deform_conv2d",
219 x.dims(),
220 weight.dims(),
221 options.weight_groups,
222 ));
223 let [batch, channels, height, width] = x.dims();
224 let [_, _, kernel_height, kernel_width] = weight.dims();
225 assert!(
226 options.offset_groups > 0 && channels.is_multiple_of(options.offset_groups),
227 "deform_conv2d input channels must be divisible by non-zero offset groups"
228 );
229 let [out_height, out_width] = options.output_size(
230 [height, width], [kernel_height, kernel_width],
231 );
232 let mask_channels = options.offset_groups * kernel_height * kernel_width;
233 assert_eq!(
234 offset.dims(), [batch, 2 * mask_channels, out_height, out_width],
235 "deform_conv2d offset shape must match groups, kernel and output"
236 );
237 if let Some(mask) = mask.as_ref() {
238 assert_eq!(
239 mask.dims(), [batch, mask_channels, out_height, out_width],
240 "deform_conv2d mask shape must match groups, kernel and output"
241 );
242 }
243 Tensor::new(TensorPrimitive::Float(B::deform_conv2d(
244 x.primitive.tensor(),
245 offset.primitive.tensor(),
246 weight.primitive.tensor(),
247 mask.map(|m| m.primitive.tensor()),
248 bias.map(|b| b.primitive.tensor()),
249 options,
250 )))
251}
252
253pub fn conv_transpose1d<B>(
255 x: Tensor<B, 3>,
256 weight: Tensor<B, 3>,
257 bias: Option<Tensor<B, 1>>,
258 options: ConvTransposeOptions<1>,
259) -> Tensor<B, 3>
260where
261 B: Backend,
262{
263 check!(TensorCheck::conv_transpose(
264 "conv_transpose1d",
265 x.dims(),
266 weight.dims(),
267 ));
268 Tensor::new(TensorPrimitive::Float(B::conv_transpose1d(
269 x.primitive.tensor(),
270 weight.primitive.tensor(),
271 bias.map(|b| b.primitive.tensor()),
272 options,
273 )))
274}
275
276pub fn conv_transpose2d<B>(
278 x: Tensor<B, 4>,
279 weight: Tensor<B, 4>,
280 bias: Option<Tensor<B, 1>>,
281 options: ConvTransposeOptions<2>,
282) -> Tensor<B, 4>
283where
284 B: Backend,
285{
286 check!(TensorCheck::conv_transpose(
287 "conv_transpose2d",
288 x.dims(),
289 weight.dims(),
290 ));
291 Tensor::new(TensorPrimitive::Float(B::conv_transpose2d(
292 x.primitive.tensor(),
293 weight.primitive.tensor(),
294 bias.map(|b| b.primitive.tensor()),
295 options,
296 )))
297}
298
299pub fn conv_transpose3d<B>(
301 x: Tensor<B, 5>,
302 weight: Tensor<B, 5>,
303 bias: Option<Tensor<B, 1>>,
304 options: ConvTransposeOptions<3>,
305) -> Tensor<B, 5>
306where
307 B: Backend,
308{
309 check!(TensorCheck::conv_transpose(
310 "conv_transpose3d",
311 x.dims(),
312 weight.dims(),
313 ));
314 Tensor::new(TensorPrimitive::Float(B::conv_transpose3d(
315 x.primitive.tensor(),
316 weight.primitive.tensor(),
317 bias.map(|b| b.primitive.tensor()),
318 options,
319 )))
320}
321
322pub fn conv_transpose1d_with_output_size<B: Backend>(
327 x: Tensor<B, 3>,
328 weight: Tensor<B, 3>,
329 bias: Option<Tensor<B, 1>>,
330 options: ConvTransposeOptions<1>,
331 output_size: usize,
332) -> Tensor<B, 3> {
333 let [_, _, input_length] = x.dims();
334 let [_, _, kernel_length] = weight.dims();
335 let options = options.with_output_size([kernel_length], [input_length], [output_size]);
336 conv_transpose1d(x, weight, bias, options)
337}
338
339pub fn conv_transpose2d_with_output_size<B: Backend>(
343 x: Tensor<B, 4>,
344 weight: Tensor<B, 4>,
345 bias: Option<Tensor<B, 1>>,
346 options: ConvTransposeOptions<2>,
347 output_size: [usize; 2],
348) -> Tensor<B, 4> {
349 let [_, _, height, width] = x.dims();
350 let [_, _, kernel_height, kernel_width] = weight.dims();
351 let options = options.with_output_size(
352 [kernel_height, kernel_width],
353 [height, width],
354 output_size,
355 );
356 conv_transpose2d(x, weight, bias, options)
357}
358
359pub fn conv_transpose3d_with_output_size<B: Backend>(
364 x: Tensor<B, 5>,
365 weight: Tensor<B, 5>,
366 bias: Option<Tensor<B, 1>>,
367 options: ConvTransposeOptions<3>,
368 output_size: [usize; 3],
369) -> Tensor<B, 5> {
370 let [_, _, depth, height, width] = x.dims();
371 let [_, _, kernel_depth, kernel_height, kernel_width] = weight.dims();
372 let options = options.with_output_size(
373 [kernel_depth, kernel_height, kernel_width],
374 [depth, height, width],
375 output_size,
376 );
377 conv_transpose3d(x, weight, bias, options)
378}
379
380pub fn unfold4d<B>(x: Tensor<B, 4>, kernel_size: [usize; 2], options: UnfoldOptions) -> Tensor<B, 3>
382where
383 B: Backend,
384{
385 Tensor::new(TensorPrimitive::Float(B::unfold4d(
386 x.primitive.tensor(),
387 kernel_size,
388 options,
389 )))
390}
391
392pub fn max_pool1d<B>(
394 x: Tensor<B, 3>,
395 kernel_size: usize,
396 stride: usize,
397 padding: usize,
398 dilation: usize,
399 ceil_mode: bool,
400) -> Tensor<B, 3>
401where
402 B: Backend,
403{
404 Tensor::new(TensorPrimitive::Float(B::max_pool1d(
405 x.primitive.tensor(),
406 kernel_size,
407 stride,
408 padding,
409 dilation,
410 ceil_mode,
411 )))
412}
413
414pub fn max_pool2d<B>(
416 x: Tensor<B, 4>,
417 kernel_size: [usize; 2],
418 stride: [usize; 2],
419 padding: [usize; 2],
420 dilation: [usize; 2],
421 ceil_mode: bool,
422) -> Tensor<B, 4>
423where
424 B: Backend,
425{
426 Tensor::new(TensorPrimitive::Float(B::max_pool2d(
427 x.primitive.tensor(),
428 kernel_size,
429 stride,
430 padding,
431 dilation,
432 ceil_mode,
433 )))
434}
435
436pub fn avg_pool2d<B>(
438 x: Tensor<B, 4>,
439 kernel_size: [usize; 2],
440 stride: [usize; 2],
441 padding: [usize; 2],
442 count_include_pad: bool,
443 ceil_mode: bool,
444) -> Tensor<B, 4>
445where
446 B: Backend,
447{
448 Tensor::new(TensorPrimitive::Float(B::avg_pool2d(
449 x.primitive.tensor(),
450 kernel_size,
451 stride,
452 padding,
453 count_include_pad,
454 ceil_mode,
455 )))
456}
457
458pub fn avg_pool1d<B>(
460 x: Tensor<B, 3>,
461 kernel_size: usize,
462 stride: usize,
463 padding: usize,
464 count_include_pad: bool,
465 ceil_mode: bool,
466) -> Tensor<B, 3>
467where
468 B: Backend,
469{
470 Tensor::new(TensorPrimitive::Float(B::avg_pool1d(
471 x.primitive.tensor(),
472 kernel_size,
473 stride,
474 padding,
475 count_include_pad,
476 ceil_mode,
477 )))
478}
479
480pub fn max_pool1d_with_indices<B>(
482 x: Tensor<B, 3>,
483 kernel_size: usize,
484 stride: usize,
485 padding: usize,
486 dilation: usize,
487 ceil_mode: bool,
488) -> (Tensor<B, 3>, Tensor<B, 3, Int>)
489where
490 B: Backend,
491{
492 let output = B::max_pool1d_with_indices(
493 x.primitive.tensor(),
494 kernel_size,
495 stride,
496 padding,
497 dilation,
498 ceil_mode,
499 );
500
501 (
502 Tensor::new(TensorPrimitive::Float(output.output)),
503 Tensor::new(output.indices),
504 )
505}
506
507pub fn max_pool2d_with_indices<B>(
509 x: Tensor<B, 4>,
510 kernel_size: [usize; 2],
511 stride: [usize; 2],
512 padding: [usize; 2],
513 dilation: [usize; 2],
514 ceil_mode: bool,
515) -> (Tensor<B, 4>, Tensor<B, 4, Int>)
516where
517 B: Backend,
518{
519 let output = B::max_pool2d_with_indices(
520 x.primitive.tensor(),
521 kernel_size,
522 stride,
523 padding,
524 dilation,
525 ceil_mode,
526 );
527
528 (
529 Tensor::new(TensorPrimitive::Float(output.output)),
530 Tensor::new(output.indices),
531 )
532}
533
534pub fn adaptive_avg_pool2d<B>(x: Tensor<B, 4>, output_size: [usize; 2]) -> Tensor<B, 4>
536where
537 B: Backend,
538{
539 Tensor::new(TensorPrimitive::Float(B::adaptive_avg_pool2d(
540 x.primitive.tensor(),
541 output_size,
542 )))
543}
544
545pub fn adaptive_avg_pool1d<B>(x: Tensor<B, 3>, output_size: usize) -> Tensor<B, 3>
547where
548 B: Backend,
549{
550 Tensor::new(TensorPrimitive::Float(B::adaptive_avg_pool1d(
551 x.primitive.tensor(),
552 output_size,
553 )))
554}
555
556pub fn interpolate<B>(
558 x: Tensor<B, 4>,
559 output_size: [usize; 2],
560 options: InterpolateOptions,
561) -> Tensor<B, 4>
562where
563 B: Backend,
564{
565 Tensor::new(TensorPrimitive::Float(B::interpolate(
566 x.primitive.tensor(),
567 output_size,
568 options,
569 )))
570}
571
572pub fn linear<B: Backend, const D: usize>(
598 input: Tensor<B, D>,
599 weight: Tensor<B, 2>,
600 bias: Option<Tensor<B, 1>>,
601) -> Tensor<B, D> {
602 if D == 1 {
603 let input = input.unsqueeze::<2>();
605 let output = linear(input, weight, bias);
606 return output.squeeze_dim(0);
607 }
608
609 Tensor::new(TensorPrimitive::Float(B::linear(
610 input.primitive.tensor(),
611 weight.primitive.tensor(),
612 bias.map(|b| b.primitive.tensor()),
613 )))
614}
615
616pub fn attention<B: Backend>(
638 query: Tensor<B, 4>,
639 key: Tensor<B, 4>,
640 value: Tensor<B, 4>,
641 mask: Option<Tensor<B, 4, Bool>>,
642 attn_bias: Option<Tensor<B, 4>>,
643 options: AttentionModuleOptions,
644) -> Tensor<B, 4> {
645 Tensor::new(TensorPrimitive::Float(B::attention(
646 query.primitive.tensor(),
647 key.primitive.tensor(),
648 value.primitive.tensor(),
649 mask.map(|mask| mask.primitive),
650 attn_bias.map(|bias| bias.primitive.tensor()),
651 options,
652 )))
653}
654
655pub fn attention_fallback<B: Backend>(
657 query: Tensor<B, 4>,
658 key: Tensor<B, 4>,
659 value: Tensor<B, 4>,
660 mask: Option<Tensor<B, 4, Bool>>,
661 attn_bias: Option<Tensor<B, 4>>,
662 options: AttentionModuleOptions,
663) -> Tensor<B, 4> {
664 Tensor::new(TensorPrimitive::Float(
665 crate::api::ops::attention::attention_fallback::<B>(
666 query.primitive.tensor(),
667 key.primitive.tensor(),
668 value.primitive.tensor(),
669 mask.map(|mask| mask.primitive),
670 attn_bias.map(|bias| bias.primitive.tensor()),
671 options,
672 ),
673 ))
674}