Skip to main content

burn_tensor/tensor/api/
pad.rs

1use alloc::vec::Vec;
2use core::ops::Range;
3
4use crate::{ElementConversion, Tensor, kind::Numeric, ops::PadMode};
5
6/// Trait for types that can be used as padding specifications.
7///
8/// Padding is specified as `(before, after)` pairs per dimension, returned as a
9/// fixed-size array `[(usize, usize); D]`. If fewer pairs than dimensions are provided,
10/// they apply to the **last** N dimensions (earlier dimensions are left unpadded).
11pub trait IntoPadding<const D: usize> {
12    /// Converts into a fixed-size array of `(before, after)` padding pairs.
13    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
33/// Backward-compatible: `(left, right, top, bottom)` maps to last 2 dimensions.
34///
35/// Equivalent to `[(top, bottom), (left, right)]`.
36impl<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
80/// Helper to build a range array for slice_assign, selecting a portion of one dimension.
81fn 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    /// Pads the tensor using the specified padding mode.
106    ///
107    /// Padding is specified as `(before, after)` pairs. If fewer pairs than tensor dimensions
108    /// are provided, they apply to the **last** N dimensions (unspecified leading dimensions
109    /// are left unpadded).
110    ///
111    /// For backward compatibility, a `(left, right, top, bottom)` tuple is also accepted,
112    /// which pads the last two dimensions.
113    ///
114    /// # Arguments
115    ///
116    /// * `padding` - Padding specification. Accepts:
117    ///   - `[(before, after); N]` fixed-size array of pairs (N <= D)
118    ///   - `&[(before, after)]` slice of pairs per dimension
119    ///   - `Vec<(before, after)>` vector of pairs
120    ///   - `(left, right, top, bottom)` tuple for last-2-dim backward compatibility
121    /// * `mode` - The padding mode: `Constant(value)`, `Reflect`, or `Edge`.
122    ///
123    /// # Returns
124    ///
125    /// A new tensor with the specified padding applied.
126    ///
127    /// # Panics
128    ///
129    /// - Panics if more padding pairs are provided than tensor dimensions.
130    /// - `Reflect` mode panics if padding exceeds `dimension_size - 1`.
131    /// - `Edge` mode panics if padding is applied to a zero-sized dimension.
132    ///
133    /// # Example
134    ///
135    /// ```rust
136    /// use burn_tensor::{Tensor, Shape};
137    /// use burn_tensor::ops::PadMode;
138    ///
139    /// let device = Default::default();
140    /// let tensor = Tensor::<2>::from_data([[12.0, -2.0, 3.0], [5.0, 3.0, 6.0]], &device);
141    ///
142    /// // Constant padding with value 0.0 (backward-compatible tuple)
143    /// let padded = tensor.clone().pad((1, 1, 1, 1), PadMode::Constant(0.0));
144    ///
145    /// // Pad arbitrary dimensions with slice of (before, after) pairs
146    /// let padded = tensor.clone().pad([(1, 1), (2, 2)], PadMode::Constant(0.0));
147    ///
148    /// // Pad only the last dimension
149    /// let padded = tensor.pad([(1, 1)], PadMode::Reflect);
150    /// ```
151    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
161/// Pad with a constant value.
162fn 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
193/// Pad using reflection at the boundaries (excluding edge values).
194///
195/// For ONNX "reflect" mode: mirrors from index 1, not index 0.
196/// Example: `[1, 2, 3, 4]` with left padding 2 becomes `[3, 2, 1, 2, 3, 4]`
197fn 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
230/// Helper to pad a single dimension using reflection.
231fn 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    // Calculate output dimensions
244    let mut output_dims = dims;
245    output_dims[dim] += pad_before + pad_after;
246
247    // Create output tensor and place original in the center
248    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    // Assign reflected "before" padding (e.g., top or left)
253    // Reflect excludes the edge, so we take indices [1..pad_before+1] and flip
254    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    // Assign reflected "after" padding (e.g., bottom or right)
262    // Take indices [dim_size - pad_after - 1..dim_size - 1] and flip
263    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
274/// Pad by replicating edge values.
275///
276/// Example: `[1, 2, 3, 4]` with left padding 2 becomes `[1, 1, 1, 2, 3, 4]`
277fn 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
304/// Helper to pad a single dimension by replicating edge values.
305fn 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    // Calculate output dimensions
318    let mut output_dims = dims;
319    output_dims[dim] += pad_before + pad_after;
320
321    // Create output tensor and place original in the center
322    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    // Assign "before" padding by repeating the first element
327    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    // Assign "after" padding by repeating the last element
335    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}