use alloc::vec;
use alloc::vec::Vec;
use burn_backend::{DType, Element};
use burn_std::{Slice, bf16, f16};
use crate::FlexTensor;
pub fn slice(tensor: FlexTensor, slices: &[Slice]) -> FlexTensor {
let (new_layout, _needs_copy) = tensor.layout().slice(slices);
FlexTensor::from_arc(tensor.data_arc(), new_layout, tensor.dtype())
}
fn normalize_index(idx: isize, dim_size: isize) -> usize {
if idx < 0 {
(dim_size + idx).max(0) as usize
} else {
idx as usize
}
}
pub fn slice_assign(tensor: FlexTensor, slices: &[Slice], value: FlexTensor) -> FlexTensor {
match tensor.dtype() {
DType::F32 => slice_assign_impl::<f32>(tensor, slices, value),
DType::F64 => slice_assign_impl::<f64>(tensor, slices, value),
DType::F16 => slice_assign_impl::<f16>(tensor, slices, value),
DType::BF16 => slice_assign_impl::<bf16>(tensor, slices, value),
DType::I32 => slice_assign_impl::<i32>(tensor, slices, value),
DType::I64 => slice_assign_impl::<i64>(tensor, slices, value),
DType::I16 => slice_assign_impl::<i16>(tensor, slices, value),
DType::I8 => slice_assign_impl::<i8>(tensor, slices, value),
DType::U32 => slice_assign_impl::<u32>(tensor, slices, value),
DType::U64 => slice_assign_impl::<u64>(tensor, slices, value),
DType::U16 => slice_assign_impl::<u16>(tensor, slices, value),
DType::U8 => slice_assign_impl::<u8>(tensor, slices, value),
DType::Bool(_) => slice_assign_impl::<u8>(tensor, slices, value),
_ => panic!("slice_assign: unsupported dtype {:?}", tensor.dtype()),
}
}
fn slice_assign_impl<E: Element + bytemuck::Pod + Clone>(
tensor: FlexTensor,
slices: &[Slice],
value: FlexTensor,
) -> FlexTensor {
if value.layout().num_elements() > 0 && value.layout().strides().iter().all(|&s| s == 0) {
let scalar = value.storage::<E>()[value.layout().start_offset()];
return slice_write_impl::<E>(tensor, slices, WriteSource::Scalar(scalar));
}
let value = value.to_contiguous();
let val_src: &[E] = value.storage::<E>();
slice_write_impl::<E>(tensor, slices, WriteSource::Slice(val_src))
}
#[derive(Copy, Clone)]
enum WriteSource<'a, E: Copy> {
Scalar(E),
Slice(&'a [E]),
}
impl<'a, E: Copy> WriteSource<'a, E> {
#[inline]
fn write_span(self, dst: &mut [E], dst_offset: usize, len: usize, src_offset: usize) {
match self {
WriteSource::Scalar(s) => dst[dst_offset..dst_offset + len].fill(s),
WriteSource::Slice(src) => dst[dst_offset..dst_offset + len]
.copy_from_slice(&src[src_offset..src_offset + len]),
}
}
#[inline]
fn write_element(self, dst: &mut [E], dst_idx: usize, src_idx: usize) {
match self {
WriteSource::Scalar(s) => dst[dst_idx] = s,
WriteSource::Slice(src) => dst[dst_idx] = src[src_idx],
}
}
}
fn slice_write_impl<E: Element + bytemuck::Pod>(
tensor: FlexTensor,
slices: &[Slice],
source: WriteSource<'_, E>,
) -> FlexTensor {
let mut tensor = if tensor.is_unique()
&& tensor.layout().is_contiguous()
&& tensor.layout().start_offset() == 0
{
tensor
} else {
tensor.into_contiguous()
};
let dst_layout = tensor.layout().clone();
let ndims = dst_layout.num_dims();
let slice_info: Vec<(usize, usize, isize)> = (0..ndims)
.map(|dim| {
let dim_size = dst_layout.shape()[dim] as isize;
let slice = if dim < slices.len() {
&slices[dim]
} else {
&Slice::new(0, None, 1)
};
compute_slice_info(slice, dim_size)
})
.collect();
let dst = tensor.storage_mut::<E>();
let inner_contiguous = slice_info
.last()
.map(|(_, _, step)| *step == 1)
.unwrap_or(false);
if ndims == 0 {
if !dst.is_empty() {
source.write_element(dst, 0, 0);
}
} else if ndims == 1 {
let (start, len, step) = slice_info[0];
if step == 1 {
source.write_span(dst, start, len, 0);
} else {
for i in 0..len {
let dst_i = if step > 0 {
start + i * step as usize
} else {
(start as isize - (i as isize) * (-step)) as usize
};
source.write_element(dst, dst_i, i);
}
}
} else if ndims == 2 && inner_contiguous {
let (row_start, row_len, row_step) = slice_info[0];
let (col_start, col_len, _) = slice_info[1];
let dst_cols = dst_layout.shape()[1];
let mut val_offset = 0;
for r in 0..row_len {
let row_idx = if row_step > 0 {
row_start + r * row_step as usize
} else {
(row_start as isize - (r as isize) * (-row_step)) as usize
};
let dst_row_start = row_idx * dst_cols + col_start;
source.write_span(dst, dst_row_start, col_len, val_offset);
val_offset += col_len;
}
} else if inner_contiguous {
let inner_len = slice_info[ndims - 1].1;
let outer_dims = ndims - 1;
let dst_strides = dst_layout.strides();
let outer_count: usize = slice_info.iter().take(outer_dims).map(|i| i.1).product();
let mut outer_indices = vec![0usize; outer_dims];
let mut val_offset = 0;
for _ in 0..outer_count {
let mut dst_offset = dst_layout.start_offset() as isize;
for (dim, &idx) in outer_indices.iter().enumerate() {
let (start, _, step) = slice_info[dim];
let src_i = if step > 0 {
start + idx * step as usize
} else {
(start as isize - (idx as isize) * (-step)) as usize
};
dst_offset += src_i as isize * dst_strides[dim];
}
dst_offset += slice_info[ndims - 1].0 as isize * dst_strides[ndims - 1];
let dst_offset = dst_offset as usize;
source.write_span(dst, dst_offset, inner_len, val_offset);
val_offset += inner_len;
for dim in (0..outer_dims).rev() {
outer_indices[dim] += 1;
if outer_indices[dim] < slice_info[dim].1 {
break;
}
outer_indices[dim] = 0;
}
}
} else {
let total_elements: usize = slice_info.iter().map(|(_, len, _)| len).product();
let dst_strides = dst_layout.strides();
let mut indices = vec![0usize; ndims];
for i in 0..total_elements {
let mut dst_offset = dst_layout.start_offset() as isize;
for (dim, &idx) in indices.iter().enumerate() {
let (start, _, step) = slice_info[dim];
let src_i = if step > 0 {
start + idx * step as usize
} else {
(start as isize - (idx as isize) * (-step)) as usize
};
dst_offset += src_i as isize * dst_strides[dim];
}
source.write_element(dst, dst_offset as usize, i);
for dim in (0..ndims).rev() {
indices[dim] += 1;
if indices[dim] < slice_info[dim].1 {
break;
}
indices[dim] = 0;
}
}
}
tensor
}
fn compute_slice_info(slice: &Slice, dim_size: isize) -> (usize, usize, isize) {
let step = slice.step;
let abs_step = step.unsigned_abs();
let range_start = normalize_index(slice.start, dim_size);
let range_end = match slice.end {
Some(e) => normalize_index(e, dim_size).min(dim_size as usize),
None => dim_size as usize,
};
let len = if range_end > range_start {
(range_end - range_start).div_ceil(abs_step)
} else {
0
};
if step > 0 {
(range_start, len, step)
} else {
let reverse_start = if range_end > range_start {
range_end - 1
} else {
range_start
};
(reverse_start, len, step)
}
}
#[cfg(test)]
mod tests {
use super::*;
use burn_backend::TensorData;
use burn_std::Shape;
#[test]
fn test_slice_basic() {
let data: Vec<f32> = vec![0.0, 1.0, 2.0, 3.0, 4.0, 5.0];
let tensor = FlexTensor::from_data(TensorData::new(data, [2, 3]));
let slices = vec![Slice::new(0, Some(1), 1), Slice::new(1, Some(3), 1)];
let result = slice(tensor, &slices);
assert_eq!(result.layout().shape().to_vec(), vec![1, 2]);
let result_data = result.into_data();
let values: Vec<f32> = bytemuck::cast_slice(&result_data.bytes).to_vec();
assert_eq!(values, vec![1.0, 2.0]);
}
#[test]
fn test_slice_with_step() {
let data: Vec<f32> = vec![0.0, 1.0, 2.0, 3.0, 4.0, 5.0];
let tensor = FlexTensor::from_data(TensorData::new(data, [6]));
let slices = vec![Slice::new(0, Some(6), 2)];
let result = slice(tensor, &slices);
assert_eq!(result.layout().shape().to_vec(), vec![3]);
let result_data = result.into_data();
let values: Vec<f32> = bytemuck::cast_slice(&result_data.bytes).to_vec();
assert_eq!(values, vec![0.0, 2.0, 4.0]);
}
#[test]
fn test_slice_negative_index() {
let data: Vec<f32> = vec![0.0, 1.0, 2.0, 3.0, 4.0];
let tensor = FlexTensor::from_data(TensorData::new(data, [5]));
let slices = vec![Slice::new(-3, None, 1)];
let result = slice(tensor, &slices);
assert_eq!(result.layout().shape().to_vec(), vec![3]);
let result_data = result.into_data();
let values: Vec<f32> = bytemuck::cast_slice(&result_data.bytes).to_vec();
assert_eq!(values, vec![2.0, 3.0, 4.0]);
}
#[test]
fn test_slice_negative_step() {
let data: Vec<f32> = vec![0.0, 1.0, 2.0, 3.0, 4.0];
let tensor = FlexTensor::from_data(TensorData::new(data, [5]));
let slices = vec![Slice::new(0, None, -1)];
let result = slice(tensor, &slices);
assert_eq!(result.layout().shape().to_vec(), vec![5]);
let result_data = result.into_data();
let values: Vec<f32> = bytemuck::cast_slice(&result_data.bytes).to_vec();
assert_eq!(values, vec![4.0, 3.0, 2.0, 1.0, 0.0]);
}
#[test]
fn test_slice_negative_step_empty_range_underflow_guard() {
let data: Vec<f32> = vec![0.0, 1.0, 2.0, 3.0, 4.0];
let tensor = FlexTensor::from_data(TensorData::new(data, [5]));
let slices1 = vec![Slice::new(0, Some(0), -1)];
let res1 = slice(tensor.clone(), &slices1);
assert_eq!(res1.layout().shape().to_vec(), vec![0]);
let slices2 = vec![Slice::new(2, Some(2), -2)];
let res2 = slice(tensor, &slices2);
assert_eq!(res2.layout().shape().to_vec(), vec![0]);
}
#[test]
fn test_slice_negative_stride_view() {
let data: Vec<f32> = vec![10.0, 20.0, 30.0, 40.0, 50.0];
let tensor = FlexTensor::from_data(TensorData::new(data, [5]));
let slices = vec![Slice::new(0, Some(5), -2)];
let res = slice(tensor, &slices);
assert_eq!(res.layout().shape().to_vec(), vec![3]);
assert_eq!(res.layout().strides(), &[-2]);
assert_eq!(res.layout().start_offset(), 4);
let values: Vec<f32> = res.into_data().try_into_vec().unwrap();
assert_eq!(values, vec![50.0, 30.0, 10.0]);
}
#[test]
fn test_slice_assign_1d() {
let data: Vec<f32> = vec![0.0, 1.0, 2.0, 3.0, 4.0];
let tensor = FlexTensor::from_data(TensorData::new(data, [5]));
let value_data: Vec<f32> = vec![10.0, 11.0, 12.0];
let value = FlexTensor::from_data(TensorData::new(value_data, [3]));
let slices = vec![Slice::new(1, Some(4), 1)];
let result = slice_assign(tensor, &slices, value);
let result_data = result.into_data();
let values: Vec<f32> = bytemuck::cast_slice(&result_data.bytes).to_vec();
assert_eq!(values, vec![0.0, 10.0, 11.0, 12.0, 4.0]);
}
#[test]
fn test_slice_assign_2d() {
let data: Vec<f32> = vec![0.0, 1.0, 2.0, 3.0, 4.0, 5.0, 6.0, 7.0, 8.0];
let tensor = FlexTensor::from_data(TensorData::new(data, [3, 3]));
let value_data: Vec<f32> = vec![10.0, 11.0, 12.0, 13.0];
let value = FlexTensor::from_data(TensorData::new(value_data, [2, 2]));
let slices = vec![Slice::new(1, Some(3), 1), Slice::new(1, Some(3), 1)];
let result = slice_assign(tensor, &slices, value);
let result_data = result.into_data();
let values: Vec<f32> = bytemuck::cast_slice(&result_data.bytes).to_vec();
assert_eq!(
values,
vec![0.0, 1.0, 2.0, 3.0, 10.0, 11.0, 6.0, 12.0, 13.0,]
);
}
#[test]
fn test_slice_assign_2d_full_row() {
let data: Vec<f32> = (0..12).map(|i| i as f32).collect();
let tensor = FlexTensor::from_data(TensorData::new(data, [3, 4]));
let value_data: Vec<f32> = vec![100.0, 101.0, 102.0, 103.0];
let value = FlexTensor::from_data(TensorData::new(value_data, [1, 4]));
let slices = vec![Slice::new(1, Some(2), 1), Slice::new(0, None, 1)];
let result = slice_assign(tensor, &slices, value);
let result_data = result.into_data();
let values: Vec<f32> = bytemuck::cast_slice(&result_data.bytes).to_vec();
assert_eq!(
values,
vec![
0.0, 1.0, 2.0, 3.0, 100.0, 101.0, 102.0, 103.0, 8.0, 9.0, 10.0, 11.0,
]
);
}
fn broadcast_scalar_f32(value: f32, target_shape: &[usize]) -> FlexTensor {
let scalar_tensor = FlexTensor::from_data(TensorData::new(vec![value], [1]));
crate::ops::expand::expand(scalar_tensor, Shape::from(target_shape.to_vec()))
}
#[test]
fn test_slice_assign_broadcast_scalar_1d_contiguous() {
let data: Vec<f32> = vec![0.0, 1.0, 2.0, 3.0, 4.0];
let tensor = FlexTensor::from_data(TensorData::new(data, [5]));
let value = broadcast_scalar_f32(7.0, &[3]);
let slices = vec![Slice::new(1, Some(4), 1)];
let result = slice_assign(tensor, &slices, value);
let result_data = result.into_data();
let values: Vec<f32> = bytemuck::cast_slice(&result_data.bytes).to_vec();
assert_eq!(values, vec![0.0, 7.0, 7.0, 7.0, 4.0]);
}
#[test]
fn test_slice_assign_broadcast_scalar_2d_inner_contiguous() {
let data: Vec<f32> = (0..16).map(|i| i as f32).collect();
let tensor = FlexTensor::from_data(TensorData::new(data, [4, 4]));
let value = broadcast_scalar_f32(-1.0, &[2, 2]);
let slices = vec![Slice::new(1, Some(3), 1), Slice::new(1, Some(3), 1)];
let result = slice_assign(tensor, &slices, value);
let result_data = result.into_data();
let values: Vec<f32> = bytemuck::cast_slice(&result_data.bytes).to_vec();
assert_eq!(
values,
vec![
0.0, 1.0, 2.0, 3.0, 4.0, -1.0, -1.0, 7.0, 8.0, -1.0, -1.0, 11.0, 12.0, 13.0, 14.0,
15.0,
]
);
}
#[test]
fn test_slice_assign_broadcast_scalar_3d_inner_contiguous() {
let data: Vec<f32> = (0..24).map(|i| i as f32).collect();
let tensor = FlexTensor::from_data(TensorData::new(data, [2, 3, 4]));
let value = broadcast_scalar_f32(9.0, &[1, 2, 2]);
let slices = vec![
Slice::new(0, Some(1), 1),
Slice::new(0, Some(2), 1),
Slice::new(1, Some(3), 1),
];
let result = slice_assign(tensor, &slices, value);
let result_data = result.into_data();
let values: Vec<f32> = bytemuck::cast_slice(&result_data.bytes).to_vec();
let mut expected: Vec<f32> = (0..24).map(|i| i as f32).collect();
for &i in &[1usize, 2, 5, 6] {
expected[i] = 9.0;
}
assert_eq!(values, expected);
}
#[test]
fn test_slice_assign_broadcast_scalar_strided_fallback() {
let data: Vec<f32> = (0..10).map(|i| i as f32).collect();
let tensor = FlexTensor::from_data(TensorData::new(data, [10]));
let value = broadcast_scalar_f32(0.0, &[5]);
let slices = vec![Slice::new(0, Some(10), 2)];
let result = slice_assign(tensor, &slices, value);
let result_data = result.into_data();
let values: Vec<f32> = bytemuck::cast_slice(&result_data.bytes).to_vec();
assert_eq!(
values,
vec![0.0, 1.0, 0.0, 3.0, 0.0, 5.0, 0.0, 7.0, 0.0, 9.0]
);
}
#[test]
fn test_slice_assign_broadcast_scalar_i64() {
fn broadcast_scalar_i64(value: i64, target_shape: &[usize]) -> FlexTensor {
let scalar_tensor = FlexTensor::from_data(TensorData::new(vec![value], [1]));
crate::ops::expand::expand(scalar_tensor, Shape::from(target_shape.to_vec()))
}
let data: Vec<i64> = (0..12).collect();
let tensor = FlexTensor::from_data(TensorData::new(data, [3, 4]));
let value = broadcast_scalar_i64(-7, &[2, 2]);
let slices = vec![Slice::new(0, Some(2), 1), Slice::new(1, Some(3), 1)];
let result = slice_assign(tensor, &slices, value);
let result_data = result.into_data();
let values: Vec<i64> = bytemuck::cast_slice(&result_data.bytes).to_vec();
assert_eq!(values, vec![0, -7, -7, 3, 4, -7, -7, 7, 8, 9, 10, 11]);
}
#[test]
fn test_slice_assign_broadcast_scalar_nd_strided_fallback() {
let data: Vec<f32> = (0..24).map(|i| i as f32).collect();
let tensor = FlexTensor::from_data(TensorData::new(data, [2, 3, 4]));
let value = broadcast_scalar_f32(9.0, &[2, 3, 2]);
let slices = vec![
Slice::new(0, Some(2), 1),
Slice::new(0, Some(3), 1),
Slice::new(0, Some(4), 2),
];
let result = slice_assign(tensor, &slices, value);
let result_data = result.into_data();
let values: Vec<f32> = bytemuck::cast_slice(&result_data.bytes).to_vec();
let mut expected: Vec<f32> = (0..24).map(|i| i as f32).collect();
for b in 0..2 {
for r in 0..3 {
for c in [0, 2] {
expected[b * 12 + r * 4 + c] = 9.0;
}
}
}
assert_eq!(values, expected);
}
#[test]
fn test_slice_assign_broadcast_scalar_2d_stepped_rows() {
let data: Vec<f32> = (0..25).map(|i| i as f32).collect();
let tensor = FlexTensor::from_data(TensorData::new(data, [5, 5]));
let value = broadcast_scalar_f32(-1.0, &[3, 3]);
let slices = vec![Slice::new(0, Some(5), 2), Slice::new(1, Some(4), 1)];
let result = slice_assign(tensor, &slices, value);
let result_data = result.into_data();
let values: Vec<f32> = bytemuck::cast_slice(&result_data.bytes).to_vec();
let mut expected: Vec<f32> = (0..25).map(|i| i as f32).collect();
for r in [0, 2, 4] {
for c in 1..4 {
expected[r * 5 + c] = -1.0;
}
}
assert_eq!(values, expected);
}
#[test]
fn test_slice_assign_unique_destination_writes_in_place() {
let tensor = FlexTensor::from_data(TensorData::new(vec![0.0f32; 8], [8]));
assert!(tensor.is_unique());
let buffer_before = tensor.bytes().as_ptr() as usize;
let value = FlexTensor::from_data(TensorData::new(vec![1.0f32, 2.0], [2]));
let result = slice_assign(tensor, &[Slice::new(0, Some(2), 1)], value);
assert_eq!(
result.bytes().as_ptr() as usize,
buffer_before,
"slice_assign copied a uniquely-owned destination instead of writing in place"
);
assert_eq!(
result.storage::<f32>(),
[1.0, 2.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0]
);
}
#[test]
fn test_slice_assign_shared_destination_copies() {
let tensor = FlexTensor::from_data(TensorData::new(vec![0.0f32; 8], [8]));
let alias = tensor.clone();
let value = FlexTensor::from_data(TensorData::new(vec![1.0f32, 2.0], [2]));
let result = slice_assign(tensor, &[Slice::new(0, Some(2), 1)], value);
assert_ne!(
result.bytes().as_ptr() as usize,
alias.bytes().as_ptr() as usize,
"slice_assign wrote through a shared destination"
);
assert_eq!(result.storage::<f32>()[..2], [1.0, 2.0]);
assert_eq!(alias.storage::<f32>(), [0.0f32; 8]);
}
#[test]
fn test_slice_assign_unique_prefix_view_writes_in_place() {
let store = FlexTensor::from_data(TensorData::new(
(0..40).map(|i| i as f32).collect::<Vec<_>>(),
[8, 5],
));
let view = store.narrow(0, 0, 5);
drop(store);
assert!(view.is_unique());
assert_eq!(view.storage::<f32>().len(), 40, "buffer outlives the view");
let buffer_before = view.bytes().as_ptr() as usize;
let value = FlexTensor::from_data(TensorData::new(vec![9.0f32; 5], [1, 5]));
let result = slice_assign(view, &[Slice::new(0, Some(1), 1)], value);
assert_eq!(
result.bytes().as_ptr() as usize,
buffer_before,
"slice_assign copied a uniquely-owned prefix view instead of writing in place"
);
let storage = result.storage::<f32>();
assert_eq!(storage[..5], [9.0; 5]);
assert_eq!(
storage[5..],
(5..40).map(|i| i as f32).collect::<Vec<_>>()[..]
);
}
#[test]
fn test_slice_assign_shared_prefix_view_compacts() {
let store = FlexTensor::from_data(TensorData::new(
(0..40).map(|i| i as f32).collect::<Vec<_>>(),
[8, 5],
));
let view = store.narrow(0, 0, 5);
assert!(!view.is_unique());
let value = FlexTensor::from_data(TensorData::new(vec![9.0f32; 5], [1, 5]));
let result = slice_assign(view, &[Slice::new(0, Some(1), 1)], value);
assert_eq!(
result.storage::<f32>().len(),
25,
"shared prefix view should be compacted to its logical size, not copied whole"
);
assert_eq!(result.storage::<f32>()[..5], [9.0; 5]);
assert_eq!(
store.storage::<f32>()[0],
0.0,
"COW did not protect the parent"
);
}
#[test]
fn test_slice_assign_unique_non_prefix_destinations_are_normalized() {
let base = || {
FlexTensor::from_data(TensorData::new(
(0..40).map(|i| i as f32).collect::<Vec<_>>(),
[8, 5],
))
};
let store = base();
let offset_view = store.narrow(0, 1, 5);
drop(store);
assert!(offset_view.is_unique() && offset_view.is_contiguous());
assert_eq!(offset_view.layout().start_offset(), 5);
let value = FlexTensor::from_data(TensorData::new(vec![9.0f32; 5], [1, 5]));
let result = slice_assign(offset_view, &[Slice::new(0, Some(1), 1)], value);
assert_eq!(
result.storage::<f32>()[..10],
[9.0, 9.0, 9.0, 9.0, 9.0, 10.0, 11.0, 12.0, 13.0, 14.0]
);
let transposed = base().transpose(0, 1);
assert!(transposed.is_unique() && !transposed.is_contiguous());
assert_eq!(transposed.layout().start_offset(), 0);
let value = FlexTensor::from_data(TensorData::new(vec![9.0f32; 8], [1, 8]));
let result = slice_assign(transposed, &[Slice::new(0, Some(1), 1)], value);
assert_eq!(
result.storage::<f32>()[..16],
[
9.0, 9.0, 9.0, 9.0, 9.0, 9.0, 9.0, 9.0, 1.0, 6.0, 11.0, 16.0, 21.0, 26.0, 31.0,
36.0
]
);
}
}