use core::mem::MaybeUninit;
use std::ops::Mul;
#[cfg(feature = "parallel")]
use smallvec::SmallVec;
use crate::view::{StridedView, StridedViewMut};
use crate::MaybeSendSync;
use crate::{broadcast_mul_into, broadcast_mul_into_uninit};
use crate::{ElementOp, Result, StridedError};
#[cfg(feature = "parallel")]
type AxisVec<T> = SmallVec<[T; 8]>;
#[cfg(not(feature = "parallel"))]
type AxisVec<T> = Vec<T>;
pub fn batched_outer_product_into<D, A, B, OpA, OpB>(
dest: &mut StridedViewMut<D>,
lhs: &StridedView<A, OpA>,
rhs: &StridedView<B, OpB>,
lhs_free_ndim: usize,
rhs_free_ndim: usize,
) -> Result<()>
where
D: Copy + MaybeSendSync + 'static,
A: Copy + MaybeSendSync + Mul<B, Output = D> + 'static,
B: Copy + MaybeSendSync + 'static,
OpA: ElementOp<A>,
OpB: ElementOp<B>,
{
validate_batched_outer_shape(dest, lhs, rhs, lhs_free_ndim, rhs_free_ndim)?;
let batch_ndim = lhs.ndim() - lhs_free_ndim;
let mut lhs_axes = AxisVec::<usize>::with_capacity(lhs.ndim());
let mut rhs_axes = AxisVec::<usize>::with_capacity(rhs.ndim());
lhs_axes.extend(0..lhs_free_ndim);
rhs_axes.extend(lhs_free_ndim..lhs_free_ndim + rhs_free_ndim);
let batch_axis_start = lhs_free_ndim + rhs_free_ndim;
lhs_axes.extend(batch_axis_start..batch_axis_start + batch_ndim);
rhs_axes.extend(batch_axis_start..batch_axis_start + batch_ndim);
broadcast_mul_into(dest, lhs, &lhs_axes, rhs, &rhs_axes)
}
pub fn batched_outer_product_into_uninit<D, A, B, OpA, OpB>(
dest: &mut StridedViewMut<MaybeUninit<D>>,
lhs: &StridedView<A, OpA>,
rhs: &StridedView<B, OpB>,
lhs_free_ndim: usize,
rhs_free_ndim: usize,
) -> Result<()>
where
D: Copy + MaybeSendSync + 'static,
A: Copy + MaybeSendSync + Mul<B, Output = D> + 'static,
B: Copy + MaybeSendSync + 'static,
OpA: ElementOp<A>,
OpB: ElementOp<B>,
{
validate_batched_outer_shape(dest, lhs, rhs, lhs_free_ndim, rhs_free_ndim)?;
let batch_ndim = lhs.ndim() - lhs_free_ndim;
let mut lhs_axes = AxisVec::<usize>::with_capacity(lhs.ndim());
let mut rhs_axes = AxisVec::<usize>::with_capacity(rhs.ndim());
lhs_axes.extend(0..lhs_free_ndim);
rhs_axes.extend(lhs_free_ndim..lhs_free_ndim + rhs_free_ndim);
let batch_axis_start = lhs_free_ndim + rhs_free_ndim;
lhs_axes.extend(batch_axis_start..batch_axis_start + batch_ndim);
rhs_axes.extend(batch_axis_start..batch_axis_start + batch_ndim);
broadcast_mul_into_uninit(dest, lhs, &lhs_axes, rhs, &rhs_axes)
}
fn validate_batched_outer_shape<D, A, OpA, B, OpB>(
dest: &StridedViewMut<D>,
lhs: &StridedView<A, OpA>,
rhs: &StridedView<B, OpB>,
lhs_free_ndim: usize,
rhs_free_ndim: usize,
) -> Result<()> {
if lhs_free_ndim > lhs.ndim() {
return Err(StridedError::RankMismatch(lhs_free_ndim, lhs.ndim()));
}
if rhs_free_ndim > rhs.ndim() {
return Err(StridedError::RankMismatch(rhs_free_ndim, rhs.ndim()));
}
let lhs_batch_ndim = lhs.ndim() - lhs_free_ndim;
let rhs_batch_ndim = rhs.ndim() - rhs_free_ndim;
if lhs_batch_ndim != rhs_batch_ndim {
return Err(StridedError::RankMismatch(lhs_batch_ndim, rhs_batch_ndim));
}
let expected_dest_rank = lhs_free_ndim + rhs_free_ndim + lhs_batch_ndim;
if dest.ndim() != expected_dest_rank {
return Err(StridedError::RankMismatch(dest.ndim(), expected_dest_rank));
}
ensure_dims(&dest.dims()[..lhs_free_ndim], &lhs.dims()[..lhs_free_ndim])?;
ensure_dims(
&dest.dims()[lhs_free_ndim..lhs_free_ndim + rhs_free_ndim],
&rhs.dims()[..rhs_free_ndim],
)?;
ensure_dims(
&dest.dims()[lhs_free_ndim + rhs_free_ndim..],
&lhs.dims()[lhs_free_ndim..],
)?;
ensure_dims(
&dest.dims()[lhs_free_ndim + rhs_free_ndim..],
&rhs.dims()[rhs_free_ndim..],
)?;
Ok(())
}
fn ensure_dims(actual: &[usize], expected: &[usize]) -> Result<()> {
if actual == expected {
Ok(())
} else {
Err(StridedError::ShapeMismatch(
actual.to_vec(),
expected.to_vec(),
))
}
}
#[derive(Clone, Debug, PartialEq, Eq)]
pub struct LazyOuterProductLayout {
pub base_dims: Vec<usize>,
pub output_strides: Vec<isize>,
}
pub fn plan_lazy_outer_product(
output_dims: &[usize],
lhs_dims: &[usize],
lhs_strides: &[isize],
lhs_axes: &[usize],
rhs_dims: &[usize],
rhs_strides: &[isize],
rhs_axes: &[usize],
) -> Result<Option<LazyOuterProductLayout>> {
for (dims, strides, axes) in [
(lhs_dims, lhs_strides, lhs_axes),
(rhs_dims, rhs_strides, rhs_axes),
] {
if dims.len() != strides.len() {
return Err(StridedError::RankMismatch(dims.len(), strides.len()));
}
if dims.len() != axes.len() {
return Err(StridedError::RankMismatch(dims.len(), axes.len()));
}
}
if !extents_match_output(output_dims, lhs_dims, lhs_axes)
|| !extents_match_output(output_dims, rhs_dims, rhs_axes)
|| lhs_strides
.iter()
.chain(rhs_strides)
.any(|&stride| stride < 0)
{
return Ok(None);
}
let Some(partition) = classify_outer_axes(output_dims.len(), lhs_axes, rhs_axes) else {
return Ok(None);
};
if checked_axes_product(lhs_dims, &partition.lhs_free)? <= 1
|| checked_axes_product(rhs_dims, &partition.rhs_free)? <= 1
{
return Ok(None);
}
let output_rank = output_dims.len();
let lhs_prefix = partition
.lhs_free_out
.iter()
.chain(&partition.rhs_free_out)
.chain(&partition.batch_out)
.copied()
.eq(0..output_rank);
let rhs_prefix = !lhs_prefix
&& partition
.rhs_free_out
.iter()
.chain(&partition.lhs_free_out)
.chain(&partition.batch_out)
.copied()
.eq(0..output_rank);
if !lhs_prefix && !rhs_prefix {
return Ok(None);
}
let lhs_physical = axes_by_physical_stride(lhs_strides, &partition.lhs_free);
let rhs_physical = axes_by_physical_stride(rhs_strides, &partition.rhs_free);
if lhs_physical == partition.lhs_free && rhs_physical == partition.rhs_free {
return Ok(None);
}
let mut base_out_axes = Vec::with_capacity(output_rank);
let (leading, leading_axes, trailing, trailing_axes) = if lhs_prefix {
(&lhs_physical, lhs_axes, &rhs_physical, rhs_axes)
} else {
(&rhs_physical, rhs_axes, &lhs_physical, lhs_axes)
};
base_out_axes.extend(leading.iter().map(|&axis| leading_axes[axis]));
base_out_axes.extend(trailing.iter().map(|&axis| trailing_axes[axis]));
base_out_axes.extend(partition.batch_out.iter().copied());
let base_dims: Vec<usize> = base_out_axes
.iter()
.map(|&axis| output_dims[axis])
.collect();
let mut output_strides = vec![0isize; output_rank];
let mut stride = 1isize;
for (&out_axis, &extent) in base_out_axes.iter().zip(&base_dims) {
output_strides[out_axis] = stride;
let extent = isize::try_from(extent).map_err(|_| StridedError::OffsetOverflow)?;
stride = stride
.checked_mul(extent)
.ok_or(StridedError::OffsetOverflow)?;
}
Ok(Some(LazyOuterProductLayout {
base_dims,
output_strides,
}))
}
struct OuterAxisPartition {
lhs_free_out: Vec<usize>,
rhs_free_out: Vec<usize>,
batch_out: Vec<usize>,
lhs_free: Vec<usize>,
rhs_free: Vec<usize>,
}
fn extents_match_output(output_dims: &[usize], dims: &[usize], axes: &[usize]) -> bool {
dims.iter()
.zip(axes)
.all(|(&dim, &axis)| output_dims.get(axis) == Some(&dim))
}
fn operand_axes_by_output(axes: &[usize], output_rank: usize) -> Option<Vec<Option<usize>>> {
let mut by_output = vec![None; output_rank];
for (operand_axis, &output_axis) in axes.iter().enumerate() {
if by_output
.get_mut(output_axis)?
.replace(operand_axis)
.is_some()
{
return None;
}
}
Some(by_output)
}
fn classify_outer_axes(
output_rank: usize,
lhs_axes: &[usize],
rhs_axes: &[usize],
) -> Option<OuterAxisPartition> {
let lhs_by_output = operand_axes_by_output(lhs_axes, output_rank)?;
let rhs_by_output = operand_axes_by_output(rhs_axes, output_rank)?;
let mut partition = OuterAxisPartition {
lhs_free_out: Vec::new(),
rhs_free_out: Vec::new(),
batch_out: Vec::new(),
lhs_free: Vec::new(),
rhs_free: Vec::new(),
};
for output_axis in 0..output_rank {
match (lhs_by_output[output_axis], rhs_by_output[output_axis]) {
(Some(_), Some(_)) => partition.batch_out.push(output_axis),
(Some(lhs_axis), None) => {
partition.lhs_free_out.push(output_axis);
partition.lhs_free.push(lhs_axis);
}
(None, Some(rhs_axis)) => {
partition.rhs_free_out.push(output_axis);
partition.rhs_free.push(rhs_axis);
}
(None, None) => return None,
}
}
Some(partition)
}
fn checked_axes_product(dims: &[usize], axes: &[usize]) -> Result<usize> {
if axes.iter().any(|&axis| dims[axis] == 0) {
return Ok(0);
}
axes.iter().try_fold(1usize, |acc, &axis| {
acc.checked_mul(dims[axis])
.ok_or(StridedError::OffsetOverflow)
})
}
fn axes_by_physical_stride(strides: &[isize], axes: &[usize]) -> Vec<usize> {
let mut sorted = axes.to_vec();
sorted.sort_by(|&lhs, &rhs| strides[lhs].cmp(&strides[rhs]).then(lhs.cmp(&rhs)));
sorted
}