Skip to main content

ruprim_host/slice/
mod.rs

1//! Slice operations for FlexTensor.
2
3use 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
11/// Slice a tensor according to the given slice parameters.
12///
13/// For positive steps, this is zero-copy (metadata only).
14/// For negative steps, data is copied to handle the reversal.
15pub fn slice(tensor: HostTensor, slices: &[Slice]) -> HostTensor {
16    let (new_layout, needs_copy) = tensor.layout().slice(slices);
17
18    if !needs_copy {
19        // Zero-copy: share data with new layout
20        HostTensor::from_arc(tensor.data_arc(), new_layout, tensor.dtype())
21    } else {
22        // Needs copy due to negative steps
23        slice_with_copy(&tensor, slices)
24    }
25}
26
27/// Slice with data copy (handles negative steps).
28fn 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
47/// Generic slice implementation with copy.
48fn 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    // Calculate output shape and collect normalized slice info
57    let mut out_shape = Vec::with_capacity(ndims);
58    let mut slice_info: Vec<(usize, usize, isize)> = Vec::with_capacity(ndims); // (start, len, step)
59
60    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            // Default: full range
67            &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    // Allocate output
84    let mut out_data: Vec<E> = Vec::with_capacity(num_elements);
85
86    // Use recursive iteration for arbitrary dimensions
87    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
94/// Recursively copy sliced elements.
95fn 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        // Base case: copy single element
107        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
120/// Compute source index from output indices and slice info.
121fn 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            // Negative step: start from high index, go down
133            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
143/// Normalize a potentially negative index to a positive one.
144fn 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
152/// Assign values to a slice of a tensor.
153pub 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
172/// Generic slice assign implementation.
173fn slice_assign_impl<E: Element + bytemuck::Pod + Clone>(
174    tensor: HostTensor,
175    slices: &[Slice],
176    value: HostTensor,
177) -> HostTensor {
178    // Broadcast-scalar fast path: if `value` is a fully-broadcast
179    // scalar (all strides zero), read the scalar once instead of
180    // materializing the expansion via `to_contiguous`. The
181    // `num_elements > 0` gate also guards the storage read against
182    // zero-sized sources where `iter().all(...)` would be vacuously
183    // true.
184    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/// Source to splat into a sliced region of a destination tensor. The
195/// two variants drive the same dispatch tree; `Scalar` hits the
196/// broadcast-scalar fast path (no value buffer), `Slice` hits the
197/// memcpy-style assign path (advances through `val_src`).
198#[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    /// Write a contiguous span of `dst` starting at `dst_offset` for
206    /// `len` elements. For `Slice`, `src_offset` is the current
207    /// position in the value buffer.
208    #[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    /// Write a single element. `src_idx` is only read in the `Slice`
218    /// variant.
219    #[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
228/// Unified slice writer used by both `slice_assign_impl` and the
229/// scalar-broadcast fast path. Walks the destination's sliced region
230/// (1D / 2D inner-contig / ND inner-contig / strided fallback) and
231/// pulls values from the given [`WriteSource`].
232fn 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        // Rank 0: single scalar destination. Only reachable from the
262        // scalar fast path; `slice_assign` on a rank-0 tensor with a
263        // rank-0 source also ends up here.
264        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            // Odometer increment over outer dims.
325            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
365/// Compute slice info (start, len, step) for a dimension.
366/// For negative step: start is the LAST index in the range (end-1), iterating down.
367fn 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    // Normalize start and end to [0, dim_size]
372    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        // Forward: start at low index, go up
386        (range_start, len, step)
387    } else {
388        // Reverse: start at end-1 (highest index in range), go down
389        // For s![2..8;-2]: start from index 7, go to 5, then 3
390        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// Tests kept here exercise flex-specific behavior: the internal
400// `slice` / `slice_assign` helpers, the broadcast-scalar fast paths for
401// `slice_fill` (1D contiguous, 2D inner-contig, 3D inner-contig, ND
402// strided fallback, stepped-row 2D inner-contig), and non-f32 dtype
403// coverage. General slice correctness across backends is covered by
404// crates/ruda-backend-tests/tests/tensor/float/ops/{slice,slice_assign}.rs.
405#[cfg(test)]
406mod tests;