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
12pub 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
45fn 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 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 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 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 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 let (grid_x, grid_y) = match padding_mode {
112 GridSamplePaddingMode::Border => {
113 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 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 (grid_x, grid_y)
127 }
128 };
129
130 let grid_x_floored = B::float_floor(grid_x.clone());
132 let grid_y_floored = B::float_floor(grid_y.clone());
133
134 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 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 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 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 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 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 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 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 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 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 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
285fn 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 let zeros = B::float_mul_scalar(coords, 0f32.into());
304 return B::float_add_scalar(zeros, min_val.into());
305 }
306
307 let period = 2.0 * span;
309
310 let x = B::float_sub_scalar(coords, min_val.into());
312 let x = B::float_abs(x);
313
314 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 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}