Skip to main content

strided_basic/
outer_product.rs

1//! Semantic outer-product API on dynamic-rank strided views.
2
3use core::mem::MaybeUninit;
4use std::ops::Mul;
5
6#[cfg(feature = "parallel")]
7use smallvec::SmallVec;
8
9use crate::view::{StridedView, StridedViewMut};
10use crate::MaybeSendSync;
11use crate::{broadcast_mul_into, broadcast_mul_into_uninit};
12use crate::{ElementOp, Result, StridedError};
13
14#[cfg(feature = "parallel")]
15type AxisVec<T> = SmallVec<[T; 8]>;
16#[cfg(not(feature = "parallel"))]
17type AxisVec<T> = Vec<T>;
18
19/// Compute `dest[lhs_free..., rhs_free..., batch...] =
20/// lhs[lhs_free..., batch...] * rhs[rhs_free..., batch...]`.
21///
22/// This is a semantic convenience wrapper over [`broadcast_mul_into`]. The
23/// broadcast/mul planner owns kernel selection, so explicit outer-product calls
24/// and equivalent broadcasted multiplication use the same implementation path.
25pub fn batched_outer_product_into<D, A, B, OpA, OpB>(
26    dest: &mut StridedViewMut<D>,
27    lhs: &StridedView<A, OpA>,
28    rhs: &StridedView<B, OpB>,
29    lhs_free_ndim: usize,
30    rhs_free_ndim: usize,
31) -> Result<()>
32where
33    D: Copy + MaybeSendSync + 'static,
34    A: Copy + MaybeSendSync + Mul<B, Output = D> + 'static,
35    B: Copy + MaybeSendSync + 'static,
36    OpA: ElementOp<A>,
37    OpB: ElementOp<B>,
38{
39    validate_batched_outer_shape(dest, lhs, rhs, lhs_free_ndim, rhs_free_ndim)?;
40
41    let batch_ndim = lhs.ndim() - lhs_free_ndim;
42    let mut lhs_axes = AxisVec::<usize>::with_capacity(lhs.ndim());
43    let mut rhs_axes = AxisVec::<usize>::with_capacity(rhs.ndim());
44
45    lhs_axes.extend(0..lhs_free_ndim);
46    rhs_axes.extend(lhs_free_ndim..lhs_free_ndim + rhs_free_ndim);
47
48    let batch_axis_start = lhs_free_ndim + rhs_free_ndim;
49    lhs_axes.extend(batch_axis_start..batch_axis_start + batch_ndim);
50    rhs_axes.extend(batch_axis_start..batch_axis_start + batch_ndim);
51
52    broadcast_mul_into(dest, lhs, &lhs_axes, rhs, &rhs_axes)
53}
54
55/// Compute a batched outer product into a fully overwritten uninitialized output.
56///
57/// Rank, shape, destination-injectivity, and reachable-byte overlap validation
58/// completes before the first write. Safe Rust borrows already prevent
59/// input/output aliasing; the explicit overlap check in the shared broadcast
60/// kernel preserves the contract for views produced through unsafe constructors.
61///
62/// `Ok(())` means every logical destination element is initialized. An error
63/// occurs before writes. A panic during replay may leave a partially initialized
64/// destination, which remains safe to drop as `MaybeUninit<D>`.
65///
66/// # Errors
67///
68/// Returns a typed rank or shape error for incompatible free/batch dimensions,
69/// [`StridedError::NonInjectiveOutputLayout`] for an overlapping output layout,
70/// [`StridedError::OverlappingInputOutput`] for aliased storage, or
71/// [`StridedError::OffsetOverflow`] when a reachable byte range is not
72/// representable.
73pub fn batched_outer_product_into_uninit<D, A, B, OpA, OpB>(
74    dest: &mut StridedViewMut<MaybeUninit<D>>,
75    lhs: &StridedView<A, OpA>,
76    rhs: &StridedView<B, OpB>,
77    lhs_free_ndim: usize,
78    rhs_free_ndim: usize,
79) -> Result<()>
80where
81    D: Copy + MaybeSendSync + 'static,
82    A: Copy + MaybeSendSync + Mul<B, Output = D> + 'static,
83    B: Copy + MaybeSendSync + 'static,
84    OpA: ElementOp<A>,
85    OpB: ElementOp<B>,
86{
87    validate_batched_outer_shape(dest, lhs, rhs, lhs_free_ndim, rhs_free_ndim)?;
88    let batch_ndim = lhs.ndim() - lhs_free_ndim;
89    let mut lhs_axes = AxisVec::<usize>::with_capacity(lhs.ndim());
90    let mut rhs_axes = AxisVec::<usize>::with_capacity(rhs.ndim());
91    lhs_axes.extend(0..lhs_free_ndim);
92    rhs_axes.extend(lhs_free_ndim..lhs_free_ndim + rhs_free_ndim);
93    let batch_axis_start = lhs_free_ndim + rhs_free_ndim;
94    lhs_axes.extend(batch_axis_start..batch_axis_start + batch_ndim);
95    rhs_axes.extend(batch_axis_start..batch_axis_start + batch_ndim);
96    broadcast_mul_into_uninit(dest, lhs, &lhs_axes, rhs, &rhs_axes)
97}
98
99fn validate_batched_outer_shape<D, A, OpA, B, OpB>(
100    dest: &StridedViewMut<D>,
101    lhs: &StridedView<A, OpA>,
102    rhs: &StridedView<B, OpB>,
103    lhs_free_ndim: usize,
104    rhs_free_ndim: usize,
105) -> Result<()> {
106    if lhs_free_ndim > lhs.ndim() {
107        return Err(StridedError::RankMismatch(lhs_free_ndim, lhs.ndim()));
108    }
109    if rhs_free_ndim > rhs.ndim() {
110        return Err(StridedError::RankMismatch(rhs_free_ndim, rhs.ndim()));
111    }
112
113    let lhs_batch_ndim = lhs.ndim() - lhs_free_ndim;
114    let rhs_batch_ndim = rhs.ndim() - rhs_free_ndim;
115    if lhs_batch_ndim != rhs_batch_ndim {
116        return Err(StridedError::RankMismatch(lhs_batch_ndim, rhs_batch_ndim));
117    }
118
119    let expected_dest_rank = lhs_free_ndim + rhs_free_ndim + lhs_batch_ndim;
120    if dest.ndim() != expected_dest_rank {
121        return Err(StridedError::RankMismatch(dest.ndim(), expected_dest_rank));
122    }
123
124    ensure_dims(&dest.dims()[..lhs_free_ndim], &lhs.dims()[..lhs_free_ndim])?;
125    ensure_dims(
126        &dest.dims()[lhs_free_ndim..lhs_free_ndim + rhs_free_ndim],
127        &rhs.dims()[..rhs_free_ndim],
128    )?;
129    ensure_dims(
130        &dest.dims()[lhs_free_ndim + rhs_free_ndim..],
131        &lhs.dims()[lhs_free_ndim..],
132    )?;
133    ensure_dims(
134        &dest.dims()[lhs_free_ndim + rhs_free_ndim..],
135        &rhs.dims()[rhs_free_ndim..],
136    )?;
137
138    Ok(())
139}
140
141fn ensure_dims(actual: &[usize], expected: &[usize]) -> Result<()> {
142    if actual == expected {
143        Ok(())
144    } else {
145        Err(StridedError::ShapeMismatch(
146            actual.to_vec(),
147            expected.to_vec(),
148        ))
149    }
150}
151
152/// Storage layout for a lazily ordered outer product.
153///
154/// Returned by [`plan_lazy_outer_product`]. The caller allocates a dense
155/// column-major base of [`base_dims`](Self::base_dims), fills it through a
156/// destination descriptor over that base with the logical output dims and
157/// [`output_strides`](Self::output_strides) at offset zero (for example with
158/// [`broadcast_mul_into_uninit`] and the original axis maps), and then exposes
159/// the result as a strided view with those same dims and strides.
160#[derive(Clone, Debug, PartialEq, Eq)]
161pub struct LazyOuterProductLayout {
162    /// Column-major extents of the dense base allocation.
163    pub base_dims: Vec<usize>,
164    /// Strides, in base elements, of the logical output axes over the base.
165    pub output_strides: Vec<isize>,
166}
167
168/// Plan an outer-product output whose memory order follows the inputs'
169/// physical stride order instead of the logical output order.
170///
171/// `lhs_axes[k]` (respectively `rhs_axes[k]`) names the output axis that
172/// operand axis `k` maps to, with the operand extent equal to the output
173/// extent. Output axes mapped only by `lhs` are its free axes, axes mapped
174/// only by `rhs` are its free axes, and axes mapped by both are batch axes.
175///
176/// A layout is returned only when the product splits into free and batch
177/// groups, both free groups have more than one element, the logical output
178/// order is `[lhs_free, rhs_free, batch]` or `[rhs_free, lhs_free, batch]`,
179/// all strides are non-negative, and sorting either operand's free axes by
180/// `(stride, axis)` changes their order. The base then holds the leading
181/// group's free axes in the leading operand's physical order, then the
182/// trailing group's free axes in the trailing operand's physical order, then
183/// the batch axes in output order. Traversing inputs in their physical order
184/// while writing the base contiguously is what makes the layout useful.
185///
186/// # Examples
187///
188/// ```
189/// use strided_basic::plan_lazy_outer_product;
190///
191/// // lhs is a transposed 2 x 3 matrix (row-major strides), rhs is a vector.
192/// let layout = plan_lazy_outer_product(&[2, 3, 4], &[2, 3], &[3, 1], &[0, 1], &[4], &[1], &[2])
193///     .unwrap()
194///     .unwrap();
195/// assert_eq!(layout.base_dims, [3, 2, 4]);
196/// assert_eq!(layout.output_strides, [3, 1, 6]);
197/// ```
198///
199/// # Errors
200///
201/// Returns [`StridedError::RankMismatch`] when an operand's dims, strides,
202/// and axis map lengths disagree, and [`StridedError::OffsetOverflow`] when a
203/// base extent product or stride is not representable. Inputs that are valid
204/// but ineligible return `Ok(None)`.
205pub fn plan_lazy_outer_product(
206    output_dims: &[usize],
207    lhs_dims: &[usize],
208    lhs_strides: &[isize],
209    lhs_axes: &[usize],
210    rhs_dims: &[usize],
211    rhs_strides: &[isize],
212    rhs_axes: &[usize],
213) -> Result<Option<LazyOuterProductLayout>> {
214    for (dims, strides, axes) in [
215        (lhs_dims, lhs_strides, lhs_axes),
216        (rhs_dims, rhs_strides, rhs_axes),
217    ] {
218        if dims.len() != strides.len() {
219            return Err(StridedError::RankMismatch(dims.len(), strides.len()));
220        }
221        if dims.len() != axes.len() {
222            return Err(StridedError::RankMismatch(dims.len(), axes.len()));
223        }
224    }
225    if !extents_match_output(output_dims, lhs_dims, lhs_axes)
226        || !extents_match_output(output_dims, rhs_dims, rhs_axes)
227        || lhs_strides
228            .iter()
229            .chain(rhs_strides)
230            .any(|&stride| stride < 0)
231    {
232        return Ok(None);
233    }
234    let Some(partition) = classify_outer_axes(output_dims.len(), lhs_axes, rhs_axes) else {
235        return Ok(None);
236    };
237    if checked_axes_product(lhs_dims, &partition.lhs_free)? <= 1
238        || checked_axes_product(rhs_dims, &partition.rhs_free)? <= 1
239    {
240        return Ok(None);
241    }
242
243    let output_rank = output_dims.len();
244    let lhs_prefix = partition
245        .lhs_free_out
246        .iter()
247        .chain(&partition.rhs_free_out)
248        .chain(&partition.batch_out)
249        .copied()
250        .eq(0..output_rank);
251    let rhs_prefix = !lhs_prefix
252        && partition
253            .rhs_free_out
254            .iter()
255            .chain(&partition.lhs_free_out)
256            .chain(&partition.batch_out)
257            .copied()
258            .eq(0..output_rank);
259    if !lhs_prefix && !rhs_prefix {
260        return Ok(None);
261    }
262
263    let lhs_physical = axes_by_physical_stride(lhs_strides, &partition.lhs_free);
264    let rhs_physical = axes_by_physical_stride(rhs_strides, &partition.rhs_free);
265    if lhs_physical == partition.lhs_free && rhs_physical == partition.rhs_free {
266        return Ok(None);
267    }
268
269    // Base axis order: leading free, trailing free, then batch in output order.
270    let mut base_out_axes = Vec::with_capacity(output_rank);
271    let (leading, leading_axes, trailing, trailing_axes) = if lhs_prefix {
272        (&lhs_physical, lhs_axes, &rhs_physical, rhs_axes)
273    } else {
274        (&rhs_physical, rhs_axes, &lhs_physical, lhs_axes)
275    };
276    base_out_axes.extend(leading.iter().map(|&axis| leading_axes[axis]));
277    base_out_axes.extend(trailing.iter().map(|&axis| trailing_axes[axis]));
278    base_out_axes.extend(partition.batch_out.iter().copied());
279
280    let base_dims: Vec<usize> = base_out_axes
281        .iter()
282        .map(|&axis| output_dims[axis])
283        .collect();
284    let mut output_strides = vec![0isize; output_rank];
285    let mut stride = 1isize;
286    for (&out_axis, &extent) in base_out_axes.iter().zip(&base_dims) {
287        // INVARIANT: `classify_outer_axes` proved every output axis appears
288        // exactly once in the free/batch groups, so each slot is written once.
289        output_strides[out_axis] = stride;
290        let extent = isize::try_from(extent).map_err(|_| StridedError::OffsetOverflow)?;
291        stride = stride
292            .checked_mul(extent)
293            .ok_or(StridedError::OffsetOverflow)?;
294    }
295    Ok(Some(LazyOuterProductLayout {
296        base_dims,
297        output_strides,
298    }))
299}
300
301struct OuterAxisPartition {
302    lhs_free_out: Vec<usize>,
303    rhs_free_out: Vec<usize>,
304    batch_out: Vec<usize>,
305    lhs_free: Vec<usize>,
306    rhs_free: Vec<usize>,
307}
308
309fn extents_match_output(output_dims: &[usize], dims: &[usize], axes: &[usize]) -> bool {
310    dims.iter()
311        .zip(axes)
312        .all(|(&dim, &axis)| output_dims.get(axis) == Some(&dim))
313}
314
315fn operand_axes_by_output(axes: &[usize], output_rank: usize) -> Option<Vec<Option<usize>>> {
316    let mut by_output = vec![None; output_rank];
317    for (operand_axis, &output_axis) in axes.iter().enumerate() {
318        if by_output
319            .get_mut(output_axis)?
320            .replace(operand_axis)
321            .is_some()
322        {
323            return None;
324        }
325    }
326    Some(by_output)
327}
328
329fn classify_outer_axes(
330    output_rank: usize,
331    lhs_axes: &[usize],
332    rhs_axes: &[usize],
333) -> Option<OuterAxisPartition> {
334    let lhs_by_output = operand_axes_by_output(lhs_axes, output_rank)?;
335    let rhs_by_output = operand_axes_by_output(rhs_axes, output_rank)?;
336    let mut partition = OuterAxisPartition {
337        lhs_free_out: Vec::new(),
338        rhs_free_out: Vec::new(),
339        batch_out: Vec::new(),
340        lhs_free: Vec::new(),
341        rhs_free: Vec::new(),
342    };
343    for output_axis in 0..output_rank {
344        match (lhs_by_output[output_axis], rhs_by_output[output_axis]) {
345            (Some(_), Some(_)) => partition.batch_out.push(output_axis),
346            (Some(lhs_axis), None) => {
347                partition.lhs_free_out.push(output_axis);
348                partition.lhs_free.push(lhs_axis);
349            }
350            (None, Some(rhs_axis)) => {
351                partition.rhs_free_out.push(output_axis);
352                partition.rhs_free.push(rhs_axis);
353            }
354            (None, None) => return None,
355        }
356    }
357    Some(partition)
358}
359
360fn checked_axes_product(dims: &[usize], axes: &[usize]) -> Result<usize> {
361    // A zero extent makes the group empty even when the other extents
362    // overflow.
363    if axes.iter().any(|&axis| dims[axis] == 0) {
364        return Ok(0);
365    }
366    axes.iter().try_fold(1usize, |acc, &axis| {
367        acc.checked_mul(dims[axis])
368            .ok_or(StridedError::OffsetOverflow)
369    })
370}
371
372fn axes_by_physical_stride(strides: &[isize], axes: &[usize]) -> Vec<usize> {
373    let mut sorted = axes.to_vec();
374    sorted.sort_by(|&lhs, &rhs| strides[lhs].cmp(&strides[rhs]).then(lhs.cmp(&rhs)));
375    sorted
376}