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    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
194/// Pad using reflection at the boundaries (excluding edge values).
195///
196/// For ONNX "reflect" mode: mirrors from index 1, not index 0.
197/// Example: `[1, 2, 3, 4]` with left padding 2 becomes `[3, 2, 1, 2, 3, 4]`
198fn 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
231/// Helper to pad a single dimension using reflection.
232fn 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    // 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, (&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    // 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    let (device, dtype) = (tensor.device(), tensor.dtype());
319
320    // Calculate output dimensions
321    let mut output_dims = dims;
322    output_dims[dim] += pad_before + pad_after;
323
324    // Create output tensor and place original in the center
325    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    // Assign "before" padding by repeating the first element
330    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    // Assign "after" padding by repeating the last element
338    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}