Skip to content

Commit 1145645

Browse files
refactor: optimize nd-dot by avoiding copy on standard layout, and iterate axis for non-contiguous
1 parent 6f2d2d0 commit 1145645

2 files changed

Lines changed: 77 additions & 32 deletions

File tree

‎src/linalg/impl_linalg.rs‎

Lines changed: 59 additions & 30 deletions
Original file line numberDiff line numberDiff line change
@@ -237,24 +237,39 @@ macro_rules! impl_dot_nd_ix2 {
237237
let k2 = rhs.shape()[0];
238238
let n = rhs.shape()[1];
239239
if k != k2 {
240-
dot_shape_error(self.len() / k, k, k2, n);
240+
panic!(
241+
"shapes {:?} and {:?} are not compatible for nd dot \
242+
(last axis of lhs must equal first axis of rhs)",
243+
self.shape(),
244+
rhs.shape()
245+
);
241246
}
242-
let rows = self.len() / k;
243-
let lhs_2d = self
244-
.to_shape((rows, k))
245-
.expect("ndarray: to_shape failed in nd dot");
246-
let result_2d = lhs_2d.dot(rhs);
247-
248-
let mut out_dim = <$dim>::zeros(ndim);
249-
for i in 0..ndim - 1 {
250-
out_dim[i] = self.shape()[i];
251-
}
252-
out_dim[ndim - 1] = n;
253247

254-
result_2d
255-
.to_shape(out_dim)
256-
.expect("ndarray: to_shape failed reshaping nd dot result")
257-
.into_owned()
248+
if self.is_standard_layout() {
249+
// C-contiguous: to_shape returns a *view* (no copy of LHS data).
250+
// Safety: rows * k == self.len() by construction.
251+
let rows = self.len() / k;
252+
let lhs_2d = self.to_shape((rows, k)).unwrap();
253+
let result_2d = lhs_2d.dot(rhs);
254+
255+
let mut out_dim = <$dim>::zeros(ndim);
256+
for i in 0..ndim - 1 {
257+
out_dim[i] = self.shape()[i];
258+
}
259+
out_dim[ndim - 1] = n;
260+
261+
// result_2d is a fresh C-contiguous owned array;
262+
// into_shape_with_order is free.
263+
result_2d.into_shape_with_order(out_dim).unwrap()
264+
} else {
265+
// Non-contiguous: iterate over the first axis so no whole-array
266+
// copy is needed. Each sub-array is (ndim-1)-D; the impl for that
267+
// dimension is already compiled (macro invocations are in order).
268+
let sub_results: Vec<_> =
269+
self.axis_iter(Axis(0)).map(|lane| lane.dot(rhs)).collect();
270+
let views: Vec<_> = sub_results.iter().map(|a| a.view()).collect();
271+
crate::stack(Axis(0), &views).unwrap()
272+
}
258273
}
259274
}
260275

@@ -280,21 +295,35 @@ where A: LinalgScalar
280295
let k2 = rhs.shape()[0];
281296
let n = rhs.shape()[1];
282297
if k != k2 {
283-
dot_shape_error(self.len() / k, k, k2, n);
298+
panic!(
299+
"shapes {:?} and {:?} are not compatible for nd dot \
300+
(last axis of lhs must equal first axis of rhs)",
301+
self.shape(),
302+
rhs.shape()
303+
);
304+
}
305+
306+
if self.is_standard_layout() {
307+
// C-contiguous: to_shape returns a *view* (no copy of LHS data).
308+
// Safety: rows * k == self.len() by construction.
309+
let rows = self.len() / k;
310+
let lhs_2d = self.to_shape((rows, k)).unwrap();
311+
let result_2d = lhs_2d.dot(rhs);
312+
313+
let mut out_shape = self.shape().to_vec();
314+
*out_shape.last_mut().unwrap() = n;
315+
316+
// result_2d is a fresh C-contiguous owned array;
317+
// into_shape_with_order is free.
318+
result_2d.into_shape_with_order(IxDyn(&out_shape)).unwrap()
319+
} else {
320+
// Non-contiguous: iterate over the first axis so no whole-array
321+
// copy is needed. Each sub-array is (ndim-1)-D IxDyn; recursion
322+
// eventually reaches contiguous 2-D which terminates the recursion.
323+
let sub_results: Vec<_> = self.axis_iter(Axis(0)).map(|lane| lane.dot(rhs)).collect();
324+
let views: Vec<_> = sub_results.iter().map(|a| a.view()).collect();
325+
crate::stack(Axis(0), &views).unwrap()
284326
}
285-
let rows = self.len() / k;
286-
let lhs_2d = self
287-
.to_shape((rows, k))
288-
.expect("ndarray: to_shape failed in nd dot (IxDyn)");
289-
let result_2d = lhs_2d.dot(rhs);
290-
291-
let mut out_shape = self.shape().to_vec();
292-
*out_shape.last_mut().unwrap() = n;
293-
294-
result_2d
295-
.to_shape(IxDyn(&out_shape))
296-
.expect("ndarray: to_shape failed reshaping nd dot result (IxDyn)")
297-
.into_owned()
298327
}
299328
}
300329

‎tests/oper.rs‎

Lines changed: 18 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -939,7 +939,7 @@ fn dot_3d_by_2d_integer()
939939
}
940940

941941
#[test]
942-
#[should_panic(expected = "not compatible for matrix multiplication")]
942+
#[should_panic(expected = "are not compatible for nd dot")]
943943
fn dot_3d_by_2d_shape_mismatch()
944944
{
945945
let lhs: Array3<f64> = Array3::zeros((3, 4, 5));
@@ -1019,7 +1019,23 @@ fn dot_dyn_5d_by_2d()
10191019
}
10201020

10211021
#[test]
1022-
#[should_panic(expected = "not compatible for matrix multiplication")]
1022+
fn dot_dyn_3d_by_2d_non_contiguous()
1023+
{
1024+
// Slice with stride 2 to get a non-contiguous layout.
1025+
let base: Array3<f64> = ArrayBuilder::new((6, 4, 5)).build();
1026+
let lhs = base.slice(s![..;2, .., ..]).into_dyn();
1027+
let rhs: Array2<f64> = ArrayBuilder::new((5, 7)).build();
1028+
1029+
let result = lhs.dot(&rhs);
1030+
assert_eq!(result.shape(), &[3, 4, 7]);
1031+
1032+
let lhs_fixed: ArrayView3<f64> = lhs.into_dimensionality::<Ix3>().unwrap();
1033+
let expected = lhs_fixed.dot(&rhs);
1034+
assert_eq!(result.into_dimensionality::<Ix3>().unwrap(), expected);
1035+
}
1036+
1037+
#[test]
1038+
#[should_panic(expected = "are not compatible for nd dot")]
10231039
fn dot_dyn_shape_mismatch()
10241040
{
10251041
let lhs: ArrayD<f64> = ArrayD::zeros(IxDyn(&[3, 4, 5]));

0 commit comments

Comments
 (0)