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    /// fn example() {
140    ///    let device = Default::default();
141    ///    let tensor = Tensor::<2>::from_data([[12.0, -2.0, 3.0], [5.0, 3.0, 6.0]], &device);
142    ///
143    ///    // Constant padding with value 0.0 (backward-compatible tuple)
144    ///    let padded = tensor.clone().pad((1, 1, 1, 1), PadMode::Constant(0.0));
145    ///
146    ///    // Pad arbitrary dimensions with slice of (before, after) pairs
147    ///    let padded = tensor.clone().pad([(1, 1), (2, 2)], PadMode::Constant(0.0));
148    ///
149    ///    // Pad only the last dimension
150    ///    let padded = tensor.pad([(1, 1)], PadMode::Reflect);
151    /// }
152    /// ```
153    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
163/// Pad with a constant value.
164fn 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
195/// Pad using reflection at the boundaries (excluding edge values).
196///
197/// For ONNX "reflect" mode: mirrors from index 1, not index 0.
198/// Example: `[1, 2, 3, 4]` with left padding 2 becomes `[3, 2, 1, 2, 3, 4]`
199fn 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
232/// Helper to pad a single dimension using reflection.
233fn 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    // Calculate output dimensions
246    let mut output_dims = dims;
247    output_dims[dim] += pad_before + pad_after;
248
249    // Create output tensor and place original in the center
250    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    // Assign reflected "before" padding (e.g., top or left)
255    // Reflect excludes the edge, so we take indices [1..pad_before+1] and flip
256    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    // Assign reflected "after" padding (e.g., bottom or right)
264    // Take indices [dim_size - pad_after - 1..dim_size - 1] and flip
265    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
276/// Pad by replicating edge values.
277///
278/// Example: `[1, 2, 3, 4]` with left padding 2 becomes `[1, 1, 1, 2, 3, 4]`
279fn 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
306/// Helper to pad a single dimension by replicating edge values.
307fn 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    // Calculate output dimensions
320    let mut output_dims = dims;
321    output_dims[dim] += pad_before + pad_after;
322
323    // Create output tensor and place original in the center
324    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    // Assign "before" padding by repeating the first element
329    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    // Assign "after" padding by repeating the last element
337    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}