1use alloc::vec::Vec;
2use core::ops::Range;
3
4use crate::{ElementConversion, Tensor, kind::Numeric, ops::PadMode};
5
6pub trait IntoPadding<const D: usize> {
12 fn into_padding(self) -> [(usize, usize); D];
14}
15
16impl<const D: usize, const N: usize> IntoPadding<D> for [(usize, usize); N] {
17 fn into_padding(self) -> [(usize, usize); D] {
18 assert!(
19 N <= D,
20 "Padding has {} pairs but tensor only has {} dimensions",
21 N,
22 D
23 );
24 let mut result = [(0usize, 0usize); D];
25 let offset = D - N;
26 for (i, pair) in self.into_iter().enumerate() {
27 result[offset + i] = pair;
28 }
29 result
30 }
31}
32
33impl<const D: usize> IntoPadding<D> for (usize, usize, usize, usize) {
37 fn into_padding(self) -> [(usize, usize); D] {
38 let (left, right, top, bottom) = self;
39 let mut result = [(0usize, 0usize); D];
40 result[D - 2] = (top, bottom);
41 result[D - 1] = (left, right);
42 result
43 }
44}
45
46impl<const D: usize> IntoPadding<D> for &[(usize, usize)] {
47 fn into_padding(self) -> [(usize, usize); D] {
48 assert!(
49 self.len() <= D,
50 "Padding has {} pairs but tensor only has {} dimensions",
51 self.len(),
52 D
53 );
54 let mut result = [(0usize, 0usize); D];
55 let offset = D - self.len();
56 for (i, &pair) in self.iter().enumerate() {
57 result[offset + i] = pair;
58 }
59 result
60 }
61}
62
63impl<const D: usize> IntoPadding<D> for Vec<(usize, usize)> {
64 fn into_padding(self) -> [(usize, usize); D] {
65 assert!(
66 self.len() <= D,
67 "Padding has {} pairs but tensor only has {} dimensions",
68 self.len(),
69 D
70 );
71 let mut result = [(0usize, 0usize); D];
72 let offset = D - self.len();
73 for (i, pair) in self.into_iter().enumerate() {
74 result[offset + i] = pair;
75 }
76 result
77 }
78}
79
80fn build_slice_ranges<const D: usize>(
82 dims: [usize; D],
83 target_dim: usize,
84 start: usize,
85 len: usize,
86) -> [Range<usize>; D] {
87 dims.iter()
88 .enumerate()
89 .map(|(i, &size)| {
90 if i == target_dim {
91 start..start + len
92 } else {
93 0..size
94 }
95 })
96 .collect::<Vec<Range<usize>>>()
97 .try_into()
98 .unwrap()
99}
100
101impl<const D: usize, K> Tensor<D, K>
102where
103 K: Numeric,
104{
105 pub fn pad(self, padding: impl IntoPadding<D>, mode: impl Into<PadMode>) -> Self {
154 let pairs = padding.into_padding();
155 match mode.into() {
156 PadMode::Constant(value) => pad_constant(self, &pairs, value),
157 PadMode::Reflect => pad_reflect(self, &pairs),
158 PadMode::Edge => pad_edge(self, &pairs),
159 }
160 }
161}
162
163fn pad_constant<const D: usize, K, E>(
165 tensor: Tensor<D, K>,
166 padding: &[(usize, usize); D],
167 value: E,
168) -> Tensor<D, K>
169where
170 K: Numeric,
171 E: ElementConversion,
172{
173 let mut padded_dims: [usize; D] = tensor.dims();
174
175 for (i, &(before, after)) in padding.iter().enumerate() {
176 padded_dims[i] += before + after;
177 }
178
179 let ranges: [Range<usize>; D] = padded_dims
180 .iter()
181 .enumerate()
182 .map(|(i, &dim)| {
183 let (before, after) = padding[i];
184 before..dim - after
185 })
186 .collect::<Vec<Range<usize>>>()
187 .try_into()
188 .unwrap();
189
190 let padded_tensor = Tensor::full(padded_dims, value, &tensor.device());
191
192 padded_tensor.slice_assign(ranges, tensor)
193}
194
195fn pad_reflect<const D: usize, K>(
200 tensor: Tensor<D, K>,
201 padding: &[(usize, usize); D],
202) -> Tensor<D, K>
203where
204 K: Numeric,
205{
206 let dims = tensor.dims();
207
208 for (i, &(before, after)) in padding.iter().enumerate() {
209 if before > 0 || after > 0 {
210 assert!(
211 before < dims[i] && after < dims[i],
212 "Reflect padding ({}, {}) must be less than dimension {} size ({})",
213 before,
214 after,
215 i,
216 dims[i]
217 );
218 }
219 }
220
221 let mut result = tensor;
222
223 for (i, &(before, after)) in padding.iter().enumerate() {
224 if before > 0 || after > 0 {
225 result = pad_reflect_dim(result, i, before, after);
226 }
227 }
228
229 result
230}
231
232fn pad_reflect_dim<const D: usize, K>(
234 tensor: Tensor<D, K>,
235 dim: usize,
236 pad_before: usize,
237 pad_after: usize,
238) -> Tensor<D, K>
239where
240 K: Numeric,
241{
242 let dims = tensor.dims();
243 let dim_size = dims[dim];
244
245 let mut output_dims = dims;
247 output_dims[dim] += pad_before + pad_after;
248
249 let output = Tensor::zeros(output_dims, &tensor.device());
251 let original_range = build_slice_ranges(output_dims, dim, pad_before, dim_size);
252 let mut output = output.slice_assign(original_range, tensor.clone());
253
254 if pad_before > 0 {
257 let before_slice = tensor.clone().narrow(dim, 1, pad_before);
258 let before_flipped = before_slice.flip([dim as isize]);
259 let before_range = build_slice_ranges(output_dims, dim, 0, pad_before);
260 output = output.slice_assign(before_range, before_flipped);
261 }
262
263 if pad_after > 0 {
266 let start = dim_size - pad_after - 1;
267 let after_slice = tensor.narrow(dim, start, pad_after);
268 let after_flipped = after_slice.flip([dim as isize]);
269 let after_range = build_slice_ranges(output_dims, dim, pad_before + dim_size, pad_after);
270 output = output.slice_assign(after_range, after_flipped);
271 }
272
273 output
274}
275
276fn pad_edge<const D: usize, K>(tensor: Tensor<D, K>, padding: &[(usize, usize); D]) -> Tensor<D, K>
280where
281 K: Numeric,
282{
283 let dims = tensor.dims();
284
285 for (i, &(before, after)) in padding.iter().enumerate() {
286 if before > 0 || after > 0 {
287 assert!(
288 dims[i] > 0,
289 "Cannot apply edge padding to zero-sized dimension {}",
290 i
291 );
292 }
293 }
294
295 let mut result = tensor;
296
297 for (i, &(before, after)) in padding.iter().enumerate() {
298 if before > 0 || after > 0 {
299 result = pad_edge_dim(result, i, before, after);
300 }
301 }
302
303 result
304}
305
306fn pad_edge_dim<const D: usize, K>(
308 tensor: Tensor<D, K>,
309 dim: usize,
310 pad_before: usize,
311 pad_after: usize,
312) -> Tensor<D, K>
313where
314 K: Numeric,
315{
316 let dims = tensor.dims();
317 let dim_size = dims[dim];
318
319 let mut output_dims = dims;
321 output_dims[dim] += pad_before + pad_after;
322
323 let output = Tensor::zeros(output_dims, &tensor.device());
325 let original_range = build_slice_ranges(output_dims, dim, pad_before, dim_size);
326 let mut output = output.slice_assign(original_range, tensor.clone());
327
328 if pad_before > 0 {
330 let first_slice = tensor.clone().narrow(dim, 0, 1);
331 let before_pad = first_slice.repeat_dim(dim, pad_before);
332 let before_range = build_slice_ranges(output_dims, dim, 0, pad_before);
333 output = output.slice_assign(before_range, before_pad);
334 }
335
336 if pad_after > 0 {
338 let last_slice = tensor.narrow(dim, dim_size - 1, 1);
339 let after_pad = last_slice.repeat_dim(dim, pad_after);
340 let after_range = build_slice_ranges(output_dims, dim, pad_before + dim_size, pad_after);
341 output = output.slice_assign(after_range, after_pad);
342 }
343
344 output
345}