Skip to main content

ruda_tensor/ops/modules/
grid_sample.rs

1use crate::{
2    Backend, TensorMetadata, get_device_settings,
3    ops::{GridSampleOptions, GridSamplePaddingMode, InterpolateMode},
4    tensor::FloatTensor,
5};
6use alloc::vec;
7use ruda_core::tensor::{Shape, Slice};
8
9mod nearest;
10mod bicubic;
11
12/// Reference implementation of grid_sample_2d that supports all options.
13///
14/// # Arguments
15///
16/// * `tensor` - The tensor being sampled from, must be contiguous with shape (N, C, H_in, W_in)
17/// * `grid` - A tensor of locations, with shape (N, H_out, W_out, 2). Values are [-1, 1].
18///   A [x = -1, y = -1] means top-left, and [x = 1, y = 1] means bottom-right
19/// * `options` - Grid sampling options
20///
21/// # Returns
22///
23/// A tensor with shape (N, C, H_out, W_out)
24pub fn float_grid_sample_2d_ref<B: Backend>(
25    tensor: FloatTensor<B>,
26    grid: FloatTensor<B>,
27    options: GridSampleOptions,
28) -> FloatTensor<B> {
29    match options.mode {
30        InterpolateMode::Nearest => nearest::sample::<B>(tensor, grid, options),
31        InterpolateMode::Bicubic => bicubic::sample::<B>(tensor, grid, options),
32        InterpolateMode::Bilinear => float_grid_sample_2d_bilinear::<B>(
33            tensor,
34            grid,
35            options.padding_mode,
36            options.align_corners,
37        ),
38        _ => todo!(
39            "Default implementation for grid_sample_2d with {:?} unimplemented",
40            options.mode
41        ),
42    }
43}
44
45/// Bilinear grid sampling implementation.
46fn float_grid_sample_2d_bilinear<B: Backend>(
47    tensor: FloatTensor<B>,
48    grid: FloatTensor<B>,
49    padding_mode: GridSamplePaddingMode,
50    align_corners: bool,
51) -> FloatTensor<B> {
52    let n = tensor.shape()[0];
53    let c = tensor.shape()[1];
54    let h_in = tensor.shape()[2];
55    let w_in = tensor.shape()[3];
56    let h_out = grid.shape()[1];
57    let w_out = grid.shape()[2];
58    let spatial_in = h_in * w_in;
59    let spatial_out = h_out * w_out;
60    let device = B::float_device(&tensor);
61
62    // Separate x and y coordinates from grid
63    // shape: (N, H_out, W_out, 1)
64    let grid_x_slice = vec![
65        Slice::new(0, Some(n as isize), 1),
66        Slice::new(0, Some(h_out as isize), 1),
67        Slice::new(0, Some(w_out as isize), 1),
68        Slice::new(0, Some(1), 1),
69    ];
70    let grid_y_slice = vec![
71        Slice::new(0, Some(n as isize), 1),
72        Slice::new(0, Some(h_out as isize), 1),
73        Slice::new(0, Some(w_out as isize), 1),
74        Slice::new(1, Some(2), 1),
75    ];
76
77    let grid_x = B::float_slice(grid.clone(), &grid_x_slice);
78    let grid_x = B::float_reshape(grid_x, Shape::new([n, 1, h_out, w_out]));
79    let grid_y = B::float_slice(grid.clone(), &grid_y_slice);
80    let grid_y = B::float_reshape(grid_y, Shape::new([n, 1, h_out, w_out]));
81
82    // Convert normalized grid coordinates [-1, 1] to pixel coordinates
83    let w_in_f = w_in as f64;
84    let h_in_f = h_in as f64;
85
86    let (grid_x, grid_y) = if align_corners {
87        // align_corners=true: x_pixel = (x_norm + 1) * (width - 1) / 2
88        // Maps -1 to 0 and 1 to width - 1
89        let grid_x = B::float_add_scalar(grid_x, 1f32.into());
90        let grid_x = B::float_mul_scalar(grid_x, ((w_in_f - 1.0) / 2.0).into());
91
92        let grid_y = B::float_add_scalar(grid_y, 1f32.into());
93        let grid_y = B::float_mul_scalar(grid_y, ((h_in_f - 1.0) / 2.0).into());
94
95        (grid_x, grid_y)
96    } else {
97        // align_corners=false: x_pixel = (x_norm + 1) * width / 2 - 0.5
98        // Maps -1 to -0.5 and 1 to width - 0.5
99        let grid_x = B::float_add_scalar(grid_x, 1f32.into());
100        let grid_x = B::float_mul_scalar(grid_x, (w_in_f / 2.0).into());
101        let grid_x = B::float_sub_scalar(grid_x, 0.5f32.into());
102
103        let grid_y = B::float_add_scalar(grid_y, 1f32.into());
104        let grid_y = B::float_mul_scalar(grid_y, (h_in_f / 2.0).into());
105        let grid_y = B::float_sub_scalar(grid_y, 0.5f32.into());
106
107        (grid_x, grid_y)
108    };
109
110    // Apply padding mode to coordinates
111    let (grid_x, grid_y) = match padding_mode {
112        GridSamplePaddingMode::Border => {
113            // Clamp coordinates to valid range [0, size-1]
114            let grid_x = B::float_clamp(grid_x, 0f32.into(), ((w_in - 1) as f32).into());
115            let grid_y = B::float_clamp(grid_y, 0f32.into(), ((h_in - 1) as f32).into());
116            (grid_x, grid_y)
117        }
118        GridSamplePaddingMode::Reflection => {
119            // Reflect coordinates at boundaries
120            let grid_x = reflect_coordinates::<B>(grid_x, w_in_f, align_corners);
121            let grid_y = reflect_coordinates::<B>(grid_y, h_in_f, align_corners);
122            (grid_x, grid_y)
123        }
124        GridSamplePaddingMode::Zeros => {
125            // Keep coordinates as-is, we'll mask out-of-bounds later
126            (grid_x, grid_y)
127        }
128    };
129
130    // Get floor indices for the four corners
131    let grid_x_floored = B::float_floor(grid_x.clone());
132    let grid_y_floored = B::float_floor(grid_y.clone());
133
134    // Compute interpolation weights (fractional part)
135    let x_frac = B::float_sub(grid_x.clone(), grid_x_floored.clone());
136    let y_frac = B::float_sub(grid_y.clone(), grid_y_floored.clone());
137
138    // Convert to integer indices
139    let settings = get_device_settings::<B>(&device);
140    let x0 = B::float_into_int(grid_x_floored.clone(), settings.int_dtype);
141    let y0 = B::float_into_int(grid_y_floored.clone(), settings.int_dtype);
142    let x1 = B::float_into_int(
143        B::float_add_scalar(grid_x_floored, 1f32.into()),
144        settings.int_dtype,
145    );
146    let y1 = B::float_into_int(
147        B::float_add_scalar(grid_y_floored, 1f32.into()),
148        settings.int_dtype,
149    );
150
151    // Create masks for out-of-bounds coordinates (only used for zeros padding)
152    let (mask_00, mask_01, mask_10, mask_11) = if padding_mode == GridSamplePaddingMode::Zeros {
153        let x0_valid = B::int_greater_equal_elem(x0.clone(), 0.into(), settings.bool_dtype);
154        let x0_valid = B::bool_and(
155            x0_valid,
156            B::int_lower_elem(x0.clone(), (w_in as i32).into(), settings.bool_dtype),
157        );
158        let x1_valid = B::int_greater_equal_elem(x1.clone(), 0.into(), settings.bool_dtype);
159        let x1_valid = B::bool_and(
160            x1_valid,
161            B::int_lower_elem(x1.clone(), (w_in as i32).into(), settings.bool_dtype),
162        );
163        let y0_valid = B::int_greater_equal_elem(y0.clone(), 0.into(), settings.bool_dtype);
164        let y0_valid = B::bool_and(
165            y0_valid,
166            B::int_lower_elem(y0.clone(), (h_in as i32).into(), settings.bool_dtype),
167        );
168        let y1_valid = B::int_greater_equal_elem(y1.clone(), 0.into(), settings.bool_dtype);
169        let y1_valid = B::bool_and(
170            y1_valid,
171            B::int_lower_elem(y1.clone(), (h_in as i32).into(), settings.bool_dtype),
172        );
173
174        (
175            Some(B::bool_and(x0_valid.clone(), y0_valid.clone())),
176            Some(B::bool_and(x0_valid.clone(), y1_valid.clone())),
177            Some(B::bool_and(x1_valid.clone(), y0_valid)),
178            Some(B::bool_and(x1_valid, y1_valid)),
179        )
180    } else {
181        (None, None, None, None)
182    };
183
184    // Clamp indices to valid range for gather
185    let x0_clamped = B::int_clamp(x0, 0.into(), ((w_in - 1) as i32).into());
186    let x1_clamped = B::int_clamp(x1, 0.into(), ((w_in - 1) as i32).into());
187    let y0_clamped = B::int_clamp(y0, 0.into(), ((h_in - 1) as i32).into());
188    let y1_clamped = B::int_clamp(y1, 0.into(), ((h_in - 1) as i32).into());
189
190    // Linear indices: idx = y * W_in + x
191    let w_in_scalar: i32 = w_in as i32;
192    let idx_00 = B::int_add(
193        B::int_mul_scalar(y0_clamped.clone(), w_in_scalar.into()),
194        x0_clamped.clone(),
195    );
196    let idx_01 = B::int_add(
197        B::int_mul_scalar(y1_clamped.clone(), w_in_scalar.into()),
198        x0_clamped,
199    );
200    let idx_10 = B::int_add(
201        B::int_mul_scalar(y0_clamped, w_in_scalar.into()),
202        x1_clamped.clone(),
203    );
204    let idx_11 = B::int_add(
205        B::int_mul_scalar(y1_clamped, w_in_scalar.into()),
206        x1_clamped,
207    );
208
209    // [N, 1, H_out, W_out] -> [N, 1, H_out * W_out]
210    let idx_00 = B::int_reshape(idx_00, Shape::new([n, 1, spatial_out]));
211    let idx_01 = B::int_reshape(idx_01, Shape::new([n, 1, spatial_out]));
212    let idx_10 = B::int_reshape(idx_10, Shape::new([n, 1, spatial_out]));
213    let idx_11 = B::int_reshape(idx_11, Shape::new([n, 1, spatial_out]));
214
215    // [N, 1, spatial] -> [N, C, spatial]
216    let idx_00 = B::int_expand(idx_00, Shape::new([n, c, spatial_out]));
217    let idx_01 = B::int_expand(idx_01, Shape::new([n, c, spatial_out]));
218    let idx_10 = B::int_expand(idx_10, Shape::new([n, c, spatial_out]));
219    let idx_11 = B::int_expand(idx_11, Shape::new([n, c, spatial_out]));
220
221    let tensor_flat = B::float_reshape(tensor, Shape::new([n, c, spatial_in]));
222
223    let sample_00 = B::float_gather(2, tensor_flat.clone(), idx_00);
224    let sample_01 = B::float_gather(2, tensor_flat.clone(), idx_01);
225    let sample_10 = B::float_gather(2, tensor_flat.clone(), idx_10);
226    let sample_11 = B::float_gather(2, tensor_flat, idx_11);
227
228    // Reshape samples to (N, C, H_out, W_out)
229    let sample_00 = B::float_reshape(sample_00, Shape::new([n, c, h_out, w_out]));
230    let sample_01 = B::float_reshape(sample_01, Shape::new([n, c, h_out, w_out]));
231    let sample_10 = B::float_reshape(sample_10, Shape::new([n, c, h_out, w_out]));
232    let sample_11 = B::float_reshape(sample_11, Shape::new([n, c, h_out, w_out]));
233
234    // Apply masks for zeros padding (set out-of-bounds samples to 0)
235    let (sample_00, sample_01, sample_10, sample_11) =
236        if padding_mode == GridSamplePaddingMode::Zeros {
237            let mask_00 = mask_00.unwrap();
238            let mask_01 = mask_01.unwrap();
239            let mask_10 = mask_10.unwrap();
240            let mask_11 = mask_11.unwrap();
241
242            let mask_00_inv = B::bool_not(mask_00);
243            let mask_00_inv = B::bool_reshape(mask_00_inv, Shape::new([n, 1, h_out, w_out]));
244            let mask_00_inv = B::bool_expand(mask_00_inv, Shape::new([n, c, h_out, w_out]));
245            let mask_01_inv = B::bool_not(mask_01);
246            let mask_01_inv = B::bool_reshape(mask_01_inv, Shape::new([n, 1, h_out, w_out]));
247            let mask_01_inv = B::bool_expand(mask_01_inv, Shape::new([n, c, h_out, w_out]));
248            let mask_10_inv = B::bool_not(mask_10);
249            let mask_10_inv = B::bool_reshape(mask_10_inv, Shape::new([n, 1, h_out, w_out]));
250            let mask_10_inv = B::bool_expand(mask_10_inv, Shape::new([n, c, h_out, w_out]));
251            let mask_11_inv = B::bool_not(mask_11);
252            let mask_11_inv = B::bool_reshape(mask_11_inv, Shape::new([n, 1, h_out, w_out]));
253            let mask_11_inv = B::bool_expand(mask_11_inv, Shape::new([n, c, h_out, w_out]));
254
255            (
256                B::float_mask_fill(sample_00, mask_00_inv, 0f32.into()),
257                B::float_mask_fill(sample_01, mask_01_inv, 0f32.into()),
258                B::float_mask_fill(sample_10, mask_10_inv, 0f32.into()),
259                B::float_mask_fill(sample_11, mask_11_inv, 0f32.into()),
260            )
261        } else {
262            (sample_00, sample_01, sample_10, sample_11)
263        };
264
265    // Compute bilinear interpolation weights
266    let one_minus_x = B::float_neg(x_frac.clone());
267    let one_minus_x = B::float_add_scalar(one_minus_x, 1f32.into());
268
269    let one_minus_y = B::float_neg(y_frac.clone());
270    let one_minus_y = B::float_add_scalar(one_minus_y, 1f32.into());
271
272    let weight_00 = B::float_mul(one_minus_x.clone(), one_minus_y.clone());
273    let weight_01 = B::float_mul(one_minus_x.clone(), y_frac.clone());
274    let weight_10 = B::float_mul(x_frac.clone(), one_minus_y);
275    let weight_11 = B::float_mul(x_frac, y_frac);
276
277    // Bilinear interpolation
278    let result = B::float_mul(sample_00, weight_00);
279    let result = B::float_add(result, B::float_mul(sample_01, weight_01));
280    let result = B::float_add(result, B::float_mul(sample_10, weight_10));
281
282    B::float_add(result, B::float_mul(sample_11, weight_11))
283}
284
285/// Reflect coordinates at boundaries using a triangle wave pattern.
286///
287/// For align_corners=true: reflects within [0, size-1]
288/// For align_corners=false: reflects within [-0.5, size-0.5]
289fn reflect_coordinates<B: Backend>(
290    coords: FloatTensor<B>,
291    size: f64,
292    align_corners: bool,
293) -> FloatTensor<B> {
294    let (min_val, max_val) = if align_corners {
295        (0.0f32, (size - 1.0) as f32)
296    } else {
297        (-0.5f32, (size - 0.5) as f32)
298    };
299
300    let span = max_val - min_val;
301    if span <= 0.0 {
302        // Edge case: size is 1, just return min_val everywhere
303        let zeros = B::float_mul_scalar(coords, 0f32.into());
304        return B::float_add_scalar(zeros, min_val.into());
305    }
306
307    // Triangle wave formula: span - |((x mod 2*span) - span)| + min_val
308    let period = 2.0 * span;
309
310    // x = abs(coord - min_val)
311    let x = B::float_sub_scalar(coords, min_val.into());
312    let x = B::float_abs(x);
313
314    // x_mod = x - floor(x / period) * period
315    let x_div = B::float_div_scalar(x.clone(), period.into());
316    let x_div_floor = B::float_floor(x_div);
317    let x_mod = B::float_sub(x, B::float_mul_scalar(x_div_floor, period.into()));
318
319    // result = span - abs(x_mod - span) + min_val
320    let diff = B::float_sub_scalar(x_mod, span.into());
321    let abs_diff = B::float_abs(diff);
322    let reflected = B::float_sub_scalar(abs_diff, span.into());
323    let reflected = B::float_neg(reflected);
324    B::float_add_scalar(reflected, min_val.into())
325}