use crate::layout::Layout;
use alloc::vec::Vec;
pub(crate) const ZIP_MAX_RANK: usize = 8;
#[derive(Debug, Clone, Copy)]
pub(crate) struct ZipNest {
pub ndim: usize,
pub shape: [usize; ZIP_MAX_RANK],
pub lhs_strides: [isize; ZIP_MAX_RANK],
pub rhs_strides: [isize; ZIP_MAX_RANK],
pub lhs_offset: usize,
pub rhs_offset: usize,
}
impl ZipNest {
#[inline]
pub fn inner(&self) -> (usize, isize, isize) {
let d = self.ndim - 1;
(self.shape[d], self.lhs_strides[d], self.rhs_strides[d])
}
pub fn lhs_is_dense_from_zero(&self) -> bool {
if self.lhs_offset != 0 {
return false;
}
let mut expected = 1isize;
for d in (0..self.ndim).rev() {
if self.lhs_strides[d] != expected {
return false;
}
match expected.checked_mul(self.shape[d] as isize) {
Some(next) => expected = next,
None => return false,
}
}
true
}
pub fn for_each_run(&self, mut f: impl FnMut(usize, usize)) {
debug_assert!(self.ndim >= 1);
let outer = self.ndim - 1;
let mut idx = [0usize; ZIP_MAX_RANK];
let mut lhs_base = self.lhs_offset as isize;
let mut rhs_base = self.rhs_offset as isize;
loop {
f(lhs_base as usize, rhs_base as usize);
let mut d = outer;
loop {
if d == 0 {
return;
}
d -= 1;
idx[d] += 1;
lhs_base += self.lhs_strides[d];
rhs_base += self.rhs_strides[d];
if idx[d] < self.shape[d] {
break;
}
idx[d] = 0;
lhs_base -= self.shape[d] as isize * self.lhs_strides[d];
rhs_base -= self.shape[d] as isize * self.rhs_strides[d];
}
}
}
}
pub(crate) fn collapse_for_zip(lhs: &Layout, rhs: &Layout) -> Option<ZipNest> {
let shape = lhs.shape();
let ndims = lhs.num_dims();
debug_assert_eq!(
&shape[..],
&rhs.shape()[..],
"collapse_for_zip: operands must be broadcast to the same shape"
);
if ndims > ZIP_MAX_RANK {
return None;
}
let lhs_strides = lhs.strides();
let rhs_strides = rhs.strides();
if lhs_strides.iter().chain(rhs_strides).any(|&s| s < 0) {
return None;
}
let mut nest = ZipNest {
ndim: 0,
shape: [0; ZIP_MAX_RANK],
lhs_strides: [0; ZIP_MAX_RANK],
rhs_strides: [0; ZIP_MAX_RANK],
lhs_offset: lhs.start_offset(),
rhs_offset: rhs.start_offset(),
};
for d in 0..ndims {
let size = shape[d];
if size == 1 {
continue;
}
let l_st = lhs_strides[d];
let r_st = rhs_strides[d];
let merge = nest.ndim > 0 && {
let prev = nest.ndim - 1;
(size as isize)
.checked_mul(l_st)
.is_some_and(|run| nest.lhs_strides[prev] == run)
&& (size as isize)
.checked_mul(r_st)
.is_some_and(|run| nest.rhs_strides[prev] == run)
};
if merge {
nest.shape[nest.ndim - 1] *= size;
nest.lhs_strides[nest.ndim - 1] = l_st;
nest.rhs_strides[nest.ndim - 1] = r_st;
} else {
nest.shape[nest.ndim] = size;
nest.lhs_strides[nest.ndim] = l_st;
nest.rhs_strides[nest.ndim] = r_st;
nest.ndim += 1;
}
}
Some(nest)
}
pub(crate) fn zip_map<E, R, F>(
lhs: &[E],
lhs_layout: &Layout,
rhs: &[E],
rhs_layout: &Layout,
op: F,
) -> Option<Vec<R>>
where
E: Copy,
F: Fn(E, E) -> R,
{
let numel = lhs_layout.num_elements();
if numel == 0 {
return Some(Vec::new());
}
let nest = collapse_for_zip(lhs_layout, rhs_layout)?;
let mut out: Vec<R> = Vec::with_capacity(numel);
if nest.ndim == 0 {
out.push(op(lhs[nest.lhs_offset], rhs[nest.rhs_offset]));
return Some(out);
}
let (len, l_st, r_st) = nest.inner();
match (l_st, r_st) {
(1, 1) => nest.for_each_run(|lb, rb| {
out.extend(
lhs[lb..lb + len]
.iter()
.zip(&rhs[rb..rb + len])
.map(|(&a, &b)| op(a, b)),
);
}),
(1, 0) => nest.for_each_run(|lb, rb| {
let b = rhs[rb];
out.extend(lhs[lb..lb + len].iter().map(|&a| op(a, b)));
}),
(0, 1) => nest.for_each_run(|lb, rb| {
let a = lhs[lb];
out.extend(rhs[rb..rb + len].iter().map(|&b| op(a, b)));
}),
_ => nest.for_each_run(|lb, rb| {
out.extend(
(0..len).map(|i| op(lhs[lb + i * l_st as usize], rhs[rb + i * r_st as usize])),
);
}),
}
debug_assert_eq!(out.len(), numel);
Some(out)
}
pub(crate) fn zip_apply_inplace<E, F>(nest: &ZipNest, dst: &mut [E], src: &[E], op: F)
where
E: Copy,
F: Fn(E, E) -> E,
{
debug_assert!(
nest.lhs_is_dense_from_zero(),
"zip_apply_inplace: destination must be dense from index 0"
);
if nest.ndim == 0 {
dst[0] = op(dst[0], src[nest.rhs_offset]);
return;
}
if nest.shape[..nest.ndim].contains(&0) {
return;
}
let (len, l_st, r_st) = nest.inner();
debug_assert_eq!(l_st, 1, "dense destination implies a contiguous inner run");
match r_st {
1 => nest.for_each_run(|lb, rb| {
for (d, &s) in dst[lb..lb + len].iter_mut().zip(&src[rb..rb + len]) {
*d = op(*d, s);
}
}),
0 => nest.for_each_run(|lb, rb| {
let s = src[rb];
for d in dst[lb..lb + len].iter_mut() {
*d = op(*d, s);
}
}),
_ => nest.for_each_run(|lb, rb| {
for (i, d) in dst[lb..lb + len].iter_mut().enumerate() {
*d = op(*d, src[rb + i * r_st as usize]);
}
}),
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::strided_index::StridedIter;
use alloc::vec;
use burn_std::Shape;
fn reference<E: Copy, R>(
lhs: &[E],
lhs_layout: &Layout,
rhs: &[E],
rhs_layout: &Layout,
op: impl Fn(E, E) -> R,
) -> Vec<R> {
StridedIter::new(lhs_layout)
.zip(StridedIter::new(rhs_layout))
.map(|(li, ri)| op(lhs[li], rhs[ri]))
.collect()
}
fn broadcast_layout(shape: &[usize], full: &[usize]) -> Layout {
let contiguous = Layout::contiguous(Shape::from(shape.to_vec()));
let mut strides = contiguous.strides().to_vec();
for (d, (&s, &f)) in shape.iter().zip(full).enumerate() {
if s == 1 && f != 1 {
strides[d] = 0;
}
}
Layout::new(Shape::from(full.to_vec()), strides, 0)
}
#[test]
fn test_collapse_contiguous_pair_merges_fully() {
let l = Layout::contiguous(Shape::from(vec![2, 3, 4]));
let r = Layout::contiguous(Shape::from(vec![2, 3, 4]));
let nest = collapse_for_zip(&l, &r).unwrap();
assert_eq!(nest.ndim, 1);
assert_eq!(nest.shape[0], 24);
assert_eq!(nest.lhs_strides[0], 1);
assert_eq!(nest.rhs_strides[0], 1);
}
#[test]
fn test_collapse_leading_broadcast_merges_inner() {
let l = Layout::contiguous(Shape::from(vec![2, 3, 4]));
let r = broadcast_layout(&[1, 3, 4], &[2, 3, 4]);
let nest = collapse_for_zip(&l, &r).unwrap();
assert_eq!(nest.ndim, 2);
assert_eq!(&nest.shape[..2], &[2, 12]);
assert_eq!(&nest.lhs_strides[..2], &[12, 1]);
assert_eq!(&nest.rhs_strides[..2], &[0, 1]);
}
#[test]
fn test_collapse_rejects_negative_strides() {
let l = Layout::contiguous(Shape::from(vec![2, 3])).flip(&[0]);
let r = Layout::contiguous(Shape::from(vec![2, 3]));
assert!(collapse_for_zip(&l, &r).is_none());
}
#[test]
fn test_zip_map_matches_strided_iter_broadcast_shapes() {
let s = 5;
let n = 7;
let full = [2usize, s, n];
let dense: Vec<f32> = (0..2 * s * n).map(|i| i as f32 * 0.5 + 1.0).collect();
let cases: Vec<(Vec<usize>, usize)> = vec![
(vec![1, s, 1], s),
(vec![1, s, n], s * n),
(vec![2, 1, 1], 2),
(vec![1, 1, 1], 1),
(vec![2, s, 1], 2 * s),
];
let dense_layout = Layout::contiguous(Shape::from(full.to_vec()));
for (bshape, belems) in cases {
let bdata: Vec<f32> = (0..belems).map(|i| i as f32 - 3.0).collect();
let blayout = broadcast_layout(&bshape, &full);
let got = zip_map(&dense, &dense_layout, &bdata, &blayout, |a, b| a * b).unwrap();
let want = reference(&dense, &dense_layout, &bdata, &blayout, |a, b| a * b);
assert_eq!(got, want, "rhs-broadcast {bshape:?}");
let got = zip_map(&bdata, &blayout, &dense, &dense_layout, |a, b| a - b).unwrap();
let want = reference(&bdata, &blayout, &dense, &dense_layout, |a, b| a - b);
assert_eq!(got, want, "lhs-broadcast {bshape:?}");
}
}
#[test]
fn test_zip_map_general_strided_inner() {
let data: Vec<i32> = (0..12).collect();
let l = Layout::contiguous(Shape::from(vec![3, 4])).transpose(0, 1); let r = Layout::contiguous(Shape::from(vec![4, 3]));
let rdata: Vec<i32> = (100..112).collect();
let got = zip_map(&data, &l, &rdata, &r, |a, b| a + b).unwrap();
let want = reference(&data, &l, &rdata, &r, |a, b| a + b);
assert_eq!(got, want);
}
#[test]
fn test_zip_map_offset_views() {
let data: Vec<f32> = (0..24).map(|i| i as f32).collect();
let l = Layout::contiguous(Shape::from(vec![4, 6])).narrow(0, 1, 2); let r = Layout::contiguous(Shape::from(vec![4, 6])).narrow(0, 2, 2); let got = zip_map(&data, &l, &data, &r, |a, b| a + b).unwrap();
let want = reference(&data, &l, &data, &r, |a, b| a + b);
assert_eq!(got, want);
}
#[test]
fn test_zip_map_single_element() {
let l = Layout::contiguous(Shape::from(vec![1, 1]));
let r = Layout::contiguous(Shape::from(vec![1, 1]));
let got = zip_map(&[3.0f32], &l, &[4.0f32], &r, |a, b| a * b).unwrap();
assert_eq!(got, vec![12.0]);
}
fn apply_inplace<E: Copy>(
dst: &mut [E],
dst_layout: &Layout,
src: &[E],
src_layout: &Layout,
op: impl Fn(E, E) -> E,
) -> Option<()> {
let nest = collapse_for_zip(dst_layout, src_layout)?;
if !nest.lhs_is_dense_from_zero() {
return None;
}
zip_apply_inplace(&nest, dst, src, op);
Some(())
}
#[test]
fn test_dense_from_zero_accepts_size_one_dim_with_stride_zero() {
let l = Layout::new(Shape::from(vec![2, 1, 4]), vec![4, 0, 1], 0);
assert!(!l.is_contiguous());
let r = broadcast_layout(&[1, 1, 4], &[2, 1, 4]);
assert!(collapse_for_zip(&l, &r).unwrap().lhs_is_dense_from_zero());
}
#[test]
fn test_dense_from_zero_rejects_broadcast_offset_and_transpose() {
let full = [2usize, 3, 4];
let dense = Layout::contiguous(Shape::from(full.to_vec()));
let bcast = broadcast_layout(&[1, 3, 4], &full);
assert!(
!collapse_for_zip(&bcast, &dense)
.unwrap()
.lhs_is_dense_from_zero()
);
let offset = Layout::contiguous(Shape::from(vec![4, 6])).narrow(0, 1, 2);
let other = Layout::contiguous(Shape::from(vec![2, 6]));
assert!(
!collapse_for_zip(&offset, &other)
.unwrap()
.lhs_is_dense_from_zero()
);
let t = Layout::contiguous(Shape::from(vec![3, 4])).transpose(0, 1);
let c = Layout::contiguous(Shape::from(vec![4, 3]));
assert!(!collapse_for_zip(&t, &c).unwrap().lhs_is_dense_from_zero());
}
#[test]
fn test_zip_apply_inplace_matches_zip_map_broadcast_shapes() {
let s = 5;
let n = 7;
let full = [2usize, s, n];
let dense: Vec<f32> = (0..2 * s * n).map(|i| i as f32 * 0.5 + 1.0).collect();
let dense_layout = Layout::contiguous(Shape::from(full.to_vec()));
let cases: Vec<(Vec<usize>, usize)> = vec![
(vec![1, s, 1], s),
(vec![1, s, n], s * n),
(vec![2, 1, 1], 2),
(vec![1, 1, 1], 1),
(vec![2, s, 1], 2 * s),
];
for (bshape, belems) in cases {
let bdata: Vec<f32> = (0..belems).map(|i| i as f32 - 3.0).collect();
let blayout = broadcast_layout(&bshape, &full);
let want = zip_map(&dense, &dense_layout, &bdata, &blayout, |a, b| a - b).unwrap();
let mut got = dense.clone();
apply_inplace(&mut got, &dense_layout, &bdata, &blayout, |d, s| d - s)
.expect("dense lhs must be reusable");
assert_eq!(got, want, "dense-lhs {bshape:?}");
let want = zip_map(&bdata, &blayout, &dense, &dense_layout, |a, b| a - b).unwrap();
let mut got = dense.clone();
apply_inplace(&mut got, &dense_layout, &bdata, &blayout, |d, s| s - d)
.expect("dense rhs must be reusable");
assert_eq!(got, want, "dense-rhs {bshape:?}");
}
}
#[test]
fn test_zip_apply_inplace_general_strided_source() {
let dst_layout = Layout::contiguous(Shape::from(vec![4, 3]));
let src_layout = Layout::contiguous(Shape::from(vec![3, 4])).transpose(0, 1);
let src: Vec<i32> = (0..12).collect();
let dst: Vec<i32> = (100..112).collect();
let want = zip_map(&dst, &dst_layout, &src, &src_layout, |a, b| a + b).unwrap();
let mut got = dst.clone();
apply_inplace(&mut got, &dst_layout, &src, &src_layout, |d, s| d + s).unwrap();
assert_eq!(got, want);
}
#[test]
fn test_zip_apply_inplace_single_element_and_empty() {
let l = Layout::contiguous(Shape::from(vec![1, 1]));
let mut dst = [3.0f32];
apply_inplace(&mut dst, &l, &[4.0f32], &l, |d, s| d * s).unwrap();
assert_eq!(dst, [12.0]);
let e = Layout::contiguous(Shape::from(vec![0, 3]));
apply_inplace::<f32>(&mut [], &e, &[], &e, |d, s| d + s).unwrap();
}
#[test]
fn test_zip_apply_inplace_leaves_trailing_storage_untouched() {
let mut data: Vec<f32> = (0..24).map(|i| i as f32).collect();
let dst_layout = Layout::new(Shape::from(vec![2, 6]), vec![6, 1], 0);
let src_layout = broadcast_layout(&[2, 1], &[2, 6]);
apply_inplace(
&mut data,
&dst_layout,
&[10.0, 20.0],
&src_layout,
|d, s| d + s,
)
.unwrap();
let mut want: Vec<f32> = (0..24).map(|i| i as f32).collect();
for (i, w) in want.iter_mut().enumerate().take(12) {
*w += if i < 6 { 10.0 } else { 20.0 };
}
assert_eq!(data, want);
}
#[test]
fn test_zip_map_empty() {
let l = Layout::contiguous(Shape::from(vec![0, 3]));
let r = Layout::contiguous(Shape::from(vec![0, 3]));
let got = zip_map::<f32, f32, _>(&[], &l, &[], &r, |a, b| a + b).unwrap();
assert!(got.is_empty());
}
}