--- a/mlx/linalg.cpp +++ b/mlx/linalg.cpp @@ -474,22 +474,23 @@ bool b_2d = b.shape(axis) == 2; auto out_type = promote_types(a.dtype(), b.dtype()); - auto ashape = a.shape(); - auto bshape = b.shape(); - - ashape[axis < 0 ? axis + a.ndim() : axis] = 3; - bshape[axis < 0 ? axis + b.ndim() : axis] = 3; + auto a_last = moveaxis(a, axis, -1, s); + auto b_last = moveaxis(b, axis, -1, s); + auto ashape = a_last.shape(); + auto bshape = b_last.shape(); + + ashape.back() = 3; + bshape.back() = 3; auto out_shape = broadcast_shapes(ashape, bshape); - if (axis < 0) { - axis += out_shape.size(); - } + auto output_axis = axis; + axis = out_shape.size() - 1; out_shape[axis] = a_2d ? 2 : 3; - auto a_ = broadcast_to(astype(a, out_type, s), out_shape, s); + auto a_ = broadcast_to(astype(a_last, out_type, s), out_shape, s); out_shape[axis] = b_2d ? 2 : 3; - auto b_ = broadcast_to(astype(b, out_type, s), out_shape, s); + auto b_ = broadcast_to(astype(b_last, out_type, s), out_shape, s); auto a_splits = split(a_, a_2d ? 2 : 3, axis); auto b_splits = split(b_, b_2d ? 2 : 3, axis); @@ -519,7 +520,7 @@ multiply(a_splits[0], b_splits[1], s), multiply(a_splits[1], b_splits[0], s), s)); - return concatenate(outputs, axis, s); + return moveaxis(concatenate(outputs, axis, s), axis, output_axis, s); } void validate_eig(