use super::*;
#[cfg(feature = "simd")]
pub(super) fn binary_op_f32<Op>(
mut lhs: HostTensor,
rhs: &HostTensor,
op: Op,
simd_hint: Option<BinaryOp>,
) -> HostTensor
where
Op: Fn(f32, f32) -> f32,
{
if simd_hint.is_some() && !lhs.layout().is_contiguous() && rhs.layout().strides().contains(&0) {
lhs = lhs.to_contiguous();
}
if let Some(simd_op) = simd_hint
&& lhs.is_unique()
&& let (Some((0, l_end)), Some((r_start, r_end))) = (
lhs.layout().contiguous_offsets(),
rhs.layout().contiguous_offsets(),
)
{
let r_slice: &[f32] = &rhs.storage()[r_start..r_end];
let lhs_storage: &mut [f32] = lhs.storage_mut();
let l_slice = &mut lhs_storage[..l_end];
match simd_op {
BinaryOp::Add => simd::add_inplace_f32(l_slice, r_slice),
BinaryOp::Sub => simd::sub_inplace_f32(l_slice, r_slice),
BinaryOp::Mul => simd::mul_inplace_f32(l_slice, r_slice),
BinaryOp::Div => simd::div_inplace_f32(l_slice, r_slice),
}
return lhs;
}
if let Some(simd_op) = simd_hint
&& let Some(pattern) = detect_broadcast_pattern(lhs.layout(), rhs)
{
return apply_broadcast_pattern_f32(lhs, rhs, simd_op, pattern);
}
binary_op_typed(lhs, rhs, op)
}
#[cfg(feature = "simd")]
#[derive(Debug, Clone, Copy)]
pub(super) enum BroadcastView {
SharedRow {
outer_count: usize,
row_len: usize,
rhs_row_offset: usize,
},
PerRowScalar {
outer_count: usize,
row_len: usize,
rhs_scalar_base: usize,
},
}
#[cfg(feature = "simd")]
pub(super) fn detect_broadcast_pattern(lhs: &Layout, rhs: &HostTensor) -> Option<BroadcastView> {
let rhs_layout = rhs.layout();
let rhs_storage_elems = rhs.storage::<f32>().len();
let (l_start, _) = lhs.contiguous_offsets()?;
if l_start != 0 {
return None;
}
let ndims = lhs.num_dims();
if ndims == 0 || rhs_layout.num_dims() != ndims {
return None;
}
let lhs_shape = lhs.shape();
let rhs_strides = rhs_layout.strides();
let last_stride = rhs_strides[ndims - 1];
if last_stride == 1 {
let outer_ok = (0..ndims - 1).all(|d| rhs_strides[d] == 0 || lhs_shape[d] == 1);
if outer_ok {
let outer_count: usize = (0..ndims - 1).map(|d| lhs_shape[d]).product();
let row_len = lhs_shape[ndims - 1];
if outer_count == 0 || row_len == 0 {
return None;
}
let rhs_row_offset = rhs_layout.start_offset();
if rhs_row_offset.checked_add(row_len)? > rhs_storage_elems {
return None;
}
return Some(BroadcastView::SharedRow {
outer_count,
row_len,
rhs_row_offset,
});
}
}
if last_stride == 0 {
let mut inner_dims = 0usize;
let mut row_len: usize = 1;
for d in (0..ndims).rev() {
if rhs_strides[d] == 0 {
inner_dims += 1;
row_len *= lhs_shape[d];
} else {
break;
}
}
if inner_dims == 0 {
return None;
}
let outer_ndims = ndims - inner_dims;
let mut expected: isize = 1;
for d in (0..outer_ndims).rev() {
if rhs_strides[d] != expected {
return None;
}
expected *= lhs_shape[d] as isize;
}
let outer_count: usize = (0..outer_ndims).map(|d| lhs_shape[d]).product();
if outer_count == 0 || row_len == 0 {
return None;
}
let rhs_scalar_base = rhs_layout.start_offset();
if rhs_scalar_base.checked_add(outer_count)? > rhs_storage_elems {
return None;
}
return Some(BroadcastView::PerRowScalar {
outer_count,
row_len,
rhs_scalar_base,
});
}
None
}
#[cfg(feature = "simd")]
pub(super) fn apply_broadcast_pattern_f32(
mut lhs: HostTensor,
rhs: &HostTensor,
simd_op: BinaryOp,
pattern: BroadcastView,
) -> HostTensor {
let numel = lhs.layout().num_elements();
let rhs_storage = rhs.storage::<f32>();
if lhs.is_unique() {
let dst = &mut lhs.storage_mut::<f32>()[..numel];
run_broadcast_pattern_f32(dst, rhs_storage, simd_op, pattern);
lhs
} else {
let mut out: Vec<f32> = lhs.storage::<f32>()[..numel].to_vec();
run_broadcast_pattern_f32(&mut out, rhs_storage, simd_op, pattern);
make_tensor(out, lhs.layout().shape().clone(), lhs.dtype())
}
}
#[cfg(feature = "simd")]
pub(super) fn run_broadcast_pattern_f32(
dst: &mut [f32],
rhs_storage: &[f32],
simd_op: BinaryOp,
pattern: BroadcastView,
) {
match pattern {
BroadcastView::SharedRow {
outer_count,
row_len,
rhs_row_offset,
} => {
let rhs_row: &[f32] = &rhs_storage[rhs_row_offset..rhs_row_offset + row_len];
let total = outer_count * row_len;
let dst_full = &mut dst[..total];
match simd_op {
BinaryOp::Add => simd::add_shared_row_inplace_f32(dst_full, rhs_row),
BinaryOp::Sub => simd::sub_shared_row_inplace_f32(dst_full, rhs_row),
BinaryOp::Mul => simd::mul_shared_row_inplace_f32(dst_full, rhs_row),
BinaryOp::Div => simd::div_shared_row_inplace_f32(dst_full, rhs_row),
}
}
BroadcastView::PerRowScalar {
outer_count,
row_len,
rhs_scalar_base,
} => {
let scalars = &rhs_storage[rhs_scalar_base..rhs_scalar_base + outer_count];
match simd_op {
BinaryOp::Add => per_row_scalar_apply(dst, scalars, row_len, |a, b| a + b),
BinaryOp::Sub => per_row_scalar_apply(dst, scalars, row_len, |a, b| a - b),
BinaryOp::Mul => per_row_scalar_apply(dst, scalars, row_len, |a, b| a * b),
BinaryOp::Div => per_row_scalar_apply(dst, scalars, row_len, |a, b| a / b),
}
}
}
}
#[cfg(feature = "simd")]
#[inline]
pub(super) fn per_row_scalar_apply<Op>(dst: &mut [f32], scalars: &[f32], row_len: usize, op: Op)
where
Op: Fn(f32, f32) -> f32,
{
for (i, &scalar) in scalars.iter().enumerate() {
let start = i * row_len;
for x in dst[start..start + row_len].iter_mut() {
*x = op(*x, scalar);
}
}
}