1use alloc::vec;
4use alloc::vec::Vec;
5use ruda_core::tensor::{DType, element::Element};
6use ruda_core::{bytes::Bytes, tensor::{Shape, Slice}};
7use half::{bf16, f16};
8
9use ruda_core::tensor::host::{HostTensor, Layout};
10
11pub fn slice(tensor: HostTensor, slices: &[Slice]) -> HostTensor {
16 let (new_layout, needs_copy) = tensor.layout().slice(slices);
17
18 if !needs_copy {
19 HostTensor::from_arc(tensor.data_arc(), new_layout, tensor.dtype())
21 } else {
22 slice_with_copy(&tensor, slices)
24 }
25}
26
27fn slice_with_copy(tensor: &HostTensor, slices: &[Slice]) -> HostTensor {
29 match tensor.dtype() {
30 DType::F32 => slice_copy_impl::<f32>(tensor, slices),
31 DType::F64 => slice_copy_impl::<f64>(tensor, slices),
32 DType::F16 => slice_copy_impl::<f16>(tensor, slices),
33 DType::BF16 => slice_copy_impl::<bf16>(tensor, slices),
34 DType::I32 => slice_copy_impl::<i32>(tensor, slices),
35 DType::I64 => slice_copy_impl::<i64>(tensor, slices),
36 DType::I16 => slice_copy_impl::<i16>(tensor, slices),
37 DType::I8 => slice_copy_impl::<i8>(tensor, slices),
38 DType::U32 => slice_copy_impl::<u32>(tensor, slices),
39 DType::U64 => slice_copy_impl::<u64>(tensor, slices),
40 DType::U16 => slice_copy_impl::<u16>(tensor, slices),
41 DType::U8 => slice_copy_impl::<u8>(tensor, slices),
42 DType::Bool(_) => slice_copy_impl::<u8>(tensor, slices),
43 _ => panic!("slice: unsupported dtype {:?}", tensor.dtype()),
44 }
45}
46
47fn slice_copy_impl<E: Element + bytemuck::Pod + Default>(
49 tensor: &HostTensor,
50 slices: &[Slice],
51) -> HostTensor {
52 let src = tensor.storage::<E>();
53 let src_layout = tensor.layout();
54 let ndims = src_layout.num_dims();
55
56 let mut out_shape = Vec::with_capacity(ndims);
58 let mut slice_info: Vec<(usize, usize, isize)> = Vec::with_capacity(ndims); for dim in 0..ndims {
61 let dim_size = src_layout.shape()[dim] as isize;
62
63 let slice = if dim < slices.len() {
64 &slices[dim]
65 } else {
66 &Slice::new(0, None, 1)
68 };
69
70 let (start, len, step) = compute_slice_info(slice, dim_size);
71 out_shape.push(len);
72 slice_info.push((start, len, step));
73 }
74
75 let out_layout = Layout::contiguous(Shape::from(out_shape.clone()));
76 let num_elements = out_layout.num_elements();
77
78 if num_elements == 0 {
79 let bytes = Bytes::from_elems::<E>(Vec::new());
80 return HostTensor::new(bytes, out_layout, tensor.dtype());
81 }
82
83 let mut out_data: Vec<E> = Vec::with_capacity(num_elements);
85
86 let mut indices = vec![0usize; ndims];
88 copy_slice_recursive(src, src_layout, &slice_info, &mut out_data, &mut indices, 0);
89
90 let bytes = Bytes::from_elems(out_data);
91 HostTensor::new(bytes, out_layout, tensor.dtype())
92}
93
94fn copy_slice_recursive<E: Copy>(
96 src: &[E],
97 src_layout: &Layout,
98 slice_info: &[(usize, usize, isize)],
99 out: &mut Vec<E>,
100 indices: &mut [usize],
101 dim: usize,
102) {
103 let ndims = src_layout.num_dims();
104
105 if dim == ndims {
106 let src_idx = compute_src_index(src_layout, slice_info, indices);
108 out.push(src[src_idx]);
109 return;
110 }
111
112 let (_, len, _) = slice_info[dim];
113
114 for i in 0..len {
115 indices[dim] = i;
116 copy_slice_recursive(src, src_layout, slice_info, out, indices, dim + 1);
117 }
118}
119
120fn compute_src_index(
122 layout: &Layout,
123 slice_info: &[(usize, usize, isize)],
124 out_indices: &[usize],
125) -> usize {
126 let mut idx = layout.start_offset() as isize;
127 for (dim, &out_i) in out_indices.iter().enumerate() {
128 let (start, _, step) = slice_info[dim];
129 let src_i = if step > 0 {
130 start + out_i * step as usize
131 } else {
132 let result = start as isize - (out_i as isize) * (-step);
134 debug_assert!(result >= 0, "slice: negative source index at dim {dim}");
135 result as usize
136 };
137 idx += src_i as isize * layout.strides()[dim];
138 }
139 debug_assert!(idx >= 0, "slice: negative final index");
140 idx as usize
141}
142
143fn normalize_index(idx: isize, dim_size: isize) -> usize {
145 if idx < 0 {
146 (dim_size + idx).max(0) as usize
147 } else {
148 idx as usize
149 }
150}
151
152pub fn slice_assign(tensor: HostTensor, slices: &[Slice], value: HostTensor) -> HostTensor {
154 match tensor.dtype() {
155 DType::F32 => slice_assign_impl::<f32>(tensor, slices, value),
156 DType::F64 => slice_assign_impl::<f64>(tensor, slices, value),
157 DType::F16 => slice_assign_impl::<f16>(tensor, slices, value),
158 DType::BF16 => slice_assign_impl::<bf16>(tensor, slices, value),
159 DType::I32 => slice_assign_impl::<i32>(tensor, slices, value),
160 DType::I64 => slice_assign_impl::<i64>(tensor, slices, value),
161 DType::I16 => slice_assign_impl::<i16>(tensor, slices, value),
162 DType::I8 => slice_assign_impl::<i8>(tensor, slices, value),
163 DType::U32 => slice_assign_impl::<u32>(tensor, slices, value),
164 DType::U64 => slice_assign_impl::<u64>(tensor, slices, value),
165 DType::U16 => slice_assign_impl::<u16>(tensor, slices, value),
166 DType::U8 => slice_assign_impl::<u8>(tensor, slices, value),
167 DType::Bool(_) => slice_assign_impl::<u8>(tensor, slices, value),
168 _ => panic!("slice_assign: unsupported dtype {:?}", tensor.dtype()),
169 }
170}
171
172fn slice_assign_impl<E: Element + bytemuck::Pod + Clone>(
174 tensor: HostTensor,
175 slices: &[Slice],
176 value: HostTensor,
177) -> HostTensor {
178 if value.layout().num_elements() > 0 && value.layout().strides().iter().all(|&s| s == 0) {
185 let scalar = value.storage::<E>()[value.layout().start_offset()];
186 return slice_write_impl::<E>(tensor, slices, WriteSource::Scalar(scalar));
187 }
188
189 let value = value.to_contiguous();
190 let val_src: &[E] = value.storage::<E>();
191 slice_write_impl::<E>(tensor, slices, WriteSource::Slice(val_src))
192}
193
194#[derive(Copy, Clone)]
199enum WriteSource<'a, E: Copy> {
200 Scalar(E),
201 Slice(&'a [E]),
202}
203
204impl<'a, E: Copy> WriteSource<'a, E> {
205 #[inline]
209 fn write_span(self, dst: &mut [E], dst_offset: usize, len: usize, src_offset: usize) {
210 match self {
211 WriteSource::Scalar(s) => dst[dst_offset..dst_offset + len].fill(s),
212 WriteSource::Slice(src) => dst[dst_offset..dst_offset + len]
213 .copy_from_slice(&src[src_offset..src_offset + len]),
214 }
215 }
216
217 #[inline]
220 fn write_element(self, dst: &mut [E], dst_idx: usize, src_idx: usize) {
221 match self {
222 WriteSource::Scalar(s) => dst[dst_idx] = s,
223 WriteSource::Slice(src) => dst[dst_idx] = src[src_idx],
224 }
225 }
226}
227
228fn slice_write_impl<E: Element + bytemuck::Pod>(
233 tensor: HostTensor,
234 slices: &[Slice],
235 source: WriteSource<'_, E>,
236) -> HostTensor {
237 let mut tensor = tensor.to_contiguous();
238 let dst_layout = tensor.layout().clone();
239 let ndims = dst_layout.num_dims();
240
241 let slice_info: Vec<(usize, usize, isize)> = (0..ndims)
242 .map(|dim| {
243 let dim_size = dst_layout.shape()[dim] as isize;
244 let slice = if dim < slices.len() {
245 &slices[dim]
246 } else {
247 &Slice::new(0, None, 1)
248 };
249 compute_slice_info(slice, dim_size)
250 })
251 .collect();
252
253 let dst = tensor.storage_mut::<E>();
254
255 let inner_contiguous = slice_info
256 .last()
257 .map(|(_, _, step)| *step == 1)
258 .unwrap_or(false);
259
260 if ndims == 0 {
261 if !dst.is_empty() {
265 source.write_element(dst, 0, 0);
266 }
267 } else if ndims == 1 {
268 let (start, len, step) = slice_info[0];
269 if step == 1 {
270 source.write_span(dst, start, len, 0);
271 } else {
272 for i in 0..len {
273 let dst_i = if step > 0 {
274 start + i * step as usize
275 } else {
276 (start as isize - (i as isize) * (-step)) as usize
277 };
278 source.write_element(dst, dst_i, i);
279 }
280 }
281 } else if ndims == 2 && inner_contiguous {
282 let (row_start, row_len, row_step) = slice_info[0];
283 let (col_start, col_len, _) = slice_info[1];
284 let dst_cols = dst_layout.shape()[1];
285
286 let mut val_offset = 0;
287 for r in 0..row_len {
288 let row_idx = if row_step > 0 {
289 row_start + r * row_step as usize
290 } else {
291 (row_start as isize - (r as isize) * (-row_step)) as usize
292 };
293 let dst_row_start = row_idx * dst_cols + col_start;
294 source.write_span(dst, dst_row_start, col_len, val_offset);
295 val_offset += col_len;
296 }
297 } else if inner_contiguous {
298 let inner_len = slice_info[ndims - 1].1;
299 let outer_dims = ndims - 1;
300 let dst_strides = dst_layout.strides();
301
302 let outer_count: usize = slice_info.iter().take(outer_dims).map(|i| i.1).product();
303
304 let mut outer_indices = vec![0usize; outer_dims];
305 let mut val_offset = 0;
306
307 for _ in 0..outer_count {
308 let mut dst_offset = dst_layout.start_offset() as isize;
309 for (dim, &idx) in outer_indices.iter().enumerate() {
310 let (start, _, step) = slice_info[dim];
311 let src_i = if step > 0 {
312 start + idx * step as usize
313 } else {
314 (start as isize - (idx as isize) * (-step)) as usize
315 };
316 dst_offset += src_i as isize * dst_strides[dim];
317 }
318 dst_offset += slice_info[ndims - 1].0 as isize * dst_strides[ndims - 1];
319 let dst_offset = dst_offset as usize;
320
321 source.write_span(dst, dst_offset, inner_len, val_offset);
322 val_offset += inner_len;
323
324 for dim in (0..outer_dims).rev() {
326 outer_indices[dim] += 1;
327 if outer_indices[dim] < slice_info[dim].1 {
328 break;
329 }
330 outer_indices[dim] = 0;
331 }
332 }
333 } else {
334 let total_elements: usize = slice_info.iter().map(|(_, len, _)| len).product();
335 let dst_strides = dst_layout.strides();
336 let mut indices = vec![0usize; ndims];
337
338 for i in 0..total_elements {
339 let mut dst_offset = dst_layout.start_offset() as isize;
340 for (dim, &idx) in indices.iter().enumerate() {
341 let (start, _, step) = slice_info[dim];
342 let src_i = if step > 0 {
343 start + idx * step as usize
344 } else {
345 (start as isize - (idx as isize) * (-step)) as usize
346 };
347 dst_offset += src_i as isize * dst_strides[dim];
348 }
349
350 source.write_element(dst, dst_offset as usize, i);
351
352 for dim in (0..ndims).rev() {
353 indices[dim] += 1;
354 if indices[dim] < slice_info[dim].1 {
355 break;
356 }
357 indices[dim] = 0;
358 }
359 }
360 }
361
362 tensor
363}
364
365fn compute_slice_info(slice: &Slice, dim_size: isize) -> (usize, usize, isize) {
368 let step = slice.step;
369 let abs_step = step.unsigned_abs();
370
371 let range_start = normalize_index(slice.start, dim_size);
373 let range_end = match slice.end {
374 Some(e) => normalize_index(e, dim_size).min(dim_size as usize),
375 None => dim_size as usize,
376 };
377
378 let len = if range_end > range_start {
379 (range_end - range_start).div_ceil(abs_step)
380 } else {
381 0
382 };
383
384 if step > 0 {
385 (range_start, len, step)
387 } else {
388 let reverse_start = if range_end > range_start {
391 range_end - 1
392 } else {
393 range_start
394 };
395 (reverse_start, len, step)
396 }
397}
398
399#[cfg(test)]
406mod tests;