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 let (device, dtype) = (tensor.device(), tensor.dtype());
173
174 for (i, &(before, after)) in padding.iter().enumerate() {
175 padded_dims[i] += before + after;
176 }
177
178 let ranges: [Range<usize>; D] = padded_dims
179 .iter()
180 .enumerate()
181 .map(|(i, &dim)| {
182 let (before, after) = padding[i];
183 before..dim - after
184 })
185 .collect::<Vec<Range<usize>>>()
186 .try_into()
187 .unwrap();
188
189 let padded_tensor = Tensor::full(padded_dims, value, (&device, dtype));
190
191 padded_tensor.slice_assign(ranges, tensor)
192}
193
194fn pad_reflect<const D: usize, K>(
199 tensor: Tensor<D, K>,
200 padding: &[(usize, usize); D],
201) -> Tensor<D, K>
202where
203 K: Numeric,
204{
205 let dims = tensor.dims();
206
207 for (i, &(before, after)) in padding.iter().enumerate() {
208 if before > 0 || after > 0 {
209 assert!(
210 before < dims[i] && after < dims[i],
211 "Reflect padding ({}, {}) must be less than dimension {} size ({})",
212 before,
213 after,
214 i,
215 dims[i]
216 );
217 }
218 }
219
220 let mut result = tensor;
221
222 for (i, &(before, after)) in padding.iter().enumerate() {
223 if before > 0 || after > 0 {
224 result = pad_reflect_dim(result, i, before, after);
225 }
226 }
227
228 result
229}
230
231fn pad_reflect_dim<const D: usize, K>(
233 tensor: Tensor<D, K>,
234 dim: usize,
235 pad_before: usize,
236 pad_after: usize,
237) -> Tensor<D, K>
238where
239 K: Numeric,
240{
241 let dims = tensor.dims();
242 let dim_size = dims[dim];
243 let (device, dtype) = (tensor.device(), tensor.dtype());
244
245 let mut output_dims = dims;
247 output_dims[dim] += pad_before + pad_after;
248
249 let output = Tensor::zeros(output_dims, (&device, dtype));
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 let (device, dtype) = (tensor.device(), tensor.dtype());
319
320 let mut output_dims = dims;
322 output_dims[dim] += pad_before + pad_after;
323
324 let output = Tensor::zeros(output_dims, (&device, dtype));
326 let original_range = build_slice_ranges(output_dims, dim, pad_before, dim_size);
327 let mut output = output.slice_assign(original_range, tensor.clone());
328
329 if pad_before > 0 {
331 let first_slice = tensor.clone().narrow(dim, 0, 1);
332 let before_pad = first_slice.repeat_dim(dim, pad_before);
333 let before_range = build_slice_ranges(output_dims, dim, 0, pad_before);
334 output = output.slice_assign(before_range, before_pad);
335 }
336
337 if pad_after > 0 {
339 let last_slice = tensor.narrow(dim, dim_size - 1, 1);
340 let after_pad = last_slice.repeat_dim(dim, pad_after);
341 let after_range = build_slice_ranges(output_dims, dim, pad_before + dim_size, pad_after);
342 output = output.slice_assign(after_range, after_pad);
343 }
344
345 output
346}