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 {
152 let pairs = padding.into_padding();
153 match mode.into() {
154 PadMode::Constant(value) => pad_constant(self, &pairs, value),
155 PadMode::Reflect => pad_reflect(self, &pairs),
156 PadMode::Edge => pad_edge(self, &pairs),
157 }
158 }
159}
160
161fn pad_constant<const D: usize, K, E>(
163 tensor: Tensor<D, K>,
164 padding: &[(usize, usize); D],
165 value: E,
166) -> Tensor<D, K>
167where
168 K: Numeric,
169 E: ElementConversion,
170{
171 let mut padded_dims: [usize; D] = tensor.dims();
172
173 for (i, &(before, after)) in padding.iter().enumerate() {
174 padded_dims[i] += before + after;
175 }
176
177 let ranges: [Range<usize>; D] = padded_dims
178 .iter()
179 .enumerate()
180 .map(|(i, &dim)| {
181 let (before, after) = padding[i];
182 before..dim - after
183 })
184 .collect::<Vec<Range<usize>>>()
185 .try_into()
186 .unwrap();
187
188 let padded_tensor = Tensor::full(padded_dims, value, &tensor.device());
189
190 padded_tensor.slice_assign(ranges, tensor)
191}
192
193fn pad_reflect<const D: usize, K>(
198 tensor: Tensor<D, K>,
199 padding: &[(usize, usize); D],
200) -> Tensor<D, K>
201where
202 K: Numeric,
203{
204 let dims = tensor.dims();
205
206 for (i, &(before, after)) in padding.iter().enumerate() {
207 if before > 0 || after > 0 {
208 assert!(
209 before < dims[i] && after < dims[i],
210 "Reflect padding ({}, {}) must be less than dimension {} size ({})",
211 before,
212 after,
213 i,
214 dims[i]
215 );
216 }
217 }
218
219 let mut result = tensor;
220
221 for (i, &(before, after)) in padding.iter().enumerate() {
222 if before > 0 || after > 0 {
223 result = pad_reflect_dim(result, i, before, after);
224 }
225 }
226
227 result
228}
229
230fn pad_reflect_dim<const D: usize, K>(
232 tensor: Tensor<D, K>,
233 dim: usize,
234 pad_before: usize,
235 pad_after: usize,
236) -> Tensor<D, K>
237where
238 K: Numeric,
239{
240 let dims = tensor.dims();
241 let dim_size = dims[dim];
242
243 let mut output_dims = dims;
245 output_dims[dim] += pad_before + pad_after;
246
247 let output = Tensor::zeros(output_dims, &tensor.device());
249 let original_range = build_slice_ranges(output_dims, dim, pad_before, dim_size);
250 let mut output = output.slice_assign(original_range, tensor.clone());
251
252 if pad_before > 0 {
255 let before_slice = tensor.clone().narrow(dim, 1, pad_before);
256 let before_flipped = before_slice.flip([dim as isize]);
257 let before_range = build_slice_ranges(output_dims, dim, 0, pad_before);
258 output = output.slice_assign(before_range, before_flipped);
259 }
260
261 if pad_after > 0 {
264 let start = dim_size - pad_after - 1;
265 let after_slice = tensor.narrow(dim, start, pad_after);
266 let after_flipped = after_slice.flip([dim as isize]);
267 let after_range = build_slice_ranges(output_dims, dim, pad_before + dim_size, pad_after);
268 output = output.slice_assign(after_range, after_flipped);
269 }
270
271 output
272}
273
274fn pad_edge<const D: usize, K>(tensor: Tensor<D, K>, padding: &[(usize, usize); D]) -> Tensor<D, K>
278where
279 K: Numeric,
280{
281 let dims = tensor.dims();
282
283 for (i, &(before, after)) in padding.iter().enumerate() {
284 if before > 0 || after > 0 {
285 assert!(
286 dims[i] > 0,
287 "Cannot apply edge padding to zero-sized dimension {}",
288 i
289 );
290 }
291 }
292
293 let mut result = tensor;
294
295 for (i, &(before, after)) in padding.iter().enumerate() {
296 if before > 0 || after > 0 {
297 result = pad_edge_dim(result, i, before, after);
298 }
299 }
300
301 result
302}
303
304fn pad_edge_dim<const D: usize, K>(
306 tensor: Tensor<D, K>,
307 dim: usize,
308 pad_before: usize,
309 pad_after: usize,
310) -> Tensor<D, K>
311where
312 K: Numeric,
313{
314 let dims = tensor.dims();
315 let dim_size = dims[dim];
316
317 let mut output_dims = dims;
319 output_dims[dim] += pad_before + pad_after;
320
321 let output = Tensor::zeros(output_dims, &tensor.device());
323 let original_range = build_slice_ranges(output_dims, dim, pad_before, dim_size);
324 let mut output = output.slice_assign(original_range, tensor.clone());
325
326 if pad_before > 0 {
328 let first_slice = tensor.clone().narrow(dim, 0, 1);
329 let before_pad = first_slice.repeat_dim(dim, pad_before);
330 let before_range = build_slice_ranges(output_dims, dim, 0, pad_before);
331 output = output.slice_assign(before_range, before_pad);
332 }
333
334 if pad_after > 0 {
336 let last_slice = tensor.narrow(dim, dim_size - 1, 1);
337 let after_pad = last_slice.repeat_dim(dim, pad_after);
338 let after_range = build_slice_ranges(output_dims, dim, pad_before + dim_size, pad_after);
339 output = output.slice_assign(after_range, after_pad);
340 }
341
342 output
343}