1use 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
19pub 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
55pub 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#[derive(Clone, Debug, PartialEq, Eq)]
161pub struct LazyOuterProductLayout {
162 pub base_dims: Vec<usize>,
164 pub output_strides: Vec<isize>,
166}
167
168pub 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 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 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 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}