# Batched Matrix Multiply

**URL:** <https://discourse.julialang.org/t/batched-matrix-multiply/42332>\
**Category:** General Usage\
**Tags:** gpu, blas, linearalgebra, cuarrays\
**Created:** [July 1, 2020, 2:11am UTC](https://discourse.julialang.org/t/batched-matrix-multiply/42332 "2020-07-01T02:11:31Z")\
**Posts on this page:** 13\
**Page:** 1

<div class="post-metadata">

**Author:** ![bmit](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/bmit/32/12443_2.png) [@bmit](https://discourse.julialang.org/u/bmit)\
**Post date:** [July 1, 2020, 2:11am UTC](https://discourse.julialang.org/t/batched-matrix-multiply/42332/1 "2020-07-01T02:11:31Z")

</div>

I’d like to be able to be able to broadcast matrix multiplication across multidimensional arrays similar to the following:

```julia
a = rand(4,3,2)
b = rand(3,4,2)

a .* b # expect a (4,4,2) array, but instead errors

```

I understand this would be ambiguous in the case of 2 4x4x2 arrays as to what I wanted to do. Is there a way currently to help broadcast out by specifying a dimension?

Now, I know I can do this in a for loop, with iteration, etc. It looks like batched matrix multiplication has already been [discussed by this community](https://discourse.julialang.org/t/optimization-based-on-intel-mkl-matrix-multiplication-batch-mode/11989). As far as I can tell this was never implemented, but I might be missing something.

The real payoff here is being able to use this syntax with some of the Array interface GPU programming provided by the CuArrays.jl/CUDA.jl packages where the parallelism can really be exploited. It looks like there is already a `gemm_batched` function wrapping the equivalent cuBLAS functionality, but I can’t access it with simple Julia Linear Algebra calls yet.

---

<div class="post-metadata">

**Author:** ![stevengj](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/stevengj/32/71_2.png) [@stevengj](https://discourse.julialang.org/u/stevengj)\
**Post date:** [July 1, 2020, 2:22am UTC](https://discourse.julialang.org/t/batched-matrix-multiply/42332/2 "2020-07-01T02:22:26Z")

</div>

You could use an array of arrays. For small matrices, using an array of `SMatrix` (from StaticArrays) should be especially efficient.

---

<div class="post-metadata">

**Author:** ![RoyiAvital](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/royiavital/32/571_2.png) [@RoyiAvital](https://discourse.julialang.org/u/RoyiAvital)\
**Post date:** [July 1, 2020, 5:24am UTC](https://discourse.julialang.org/t/batched-matrix-multiply/42332/3 "2020-07-01T05:24:19Z")

</div>

There are so many features in `MKL` that can improve many real world use cases.  
I wish 2 things happened:

1. The integration of `MKL` (Be it `MKL.jl` will take advantage of that).
2. Julia will work with OpenBLAS to implement them as well.

---

<div class="post-metadata">

**Author:** ![mcabbott](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/mcabbott/32/6603_2.png) [@mcabbott](https://discourse.julialang.org/u/mcabbott)\
**Post date:** [July 1, 2020, 7:01am UTC](https://discourse.julialang.org/t/batched-matrix-multiply/42332/4 "2020-07-01T07:01:25Z")

</div>

You are probably looking for `NNlib.batched_mul`.

On the CPU this is a simple loop, because nobody has got around to hooking it up to the special MKL routines. On the GPU it calls the cuBLAS function.

---

<div class="post-metadata">

**Author:** ![DNF](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/dnf/32/10191_2.png) [@DNF](https://discourse.julialang.org/u/DNF)\
**Post date:** [July 1, 2020, 7:53am UTC](https://discourse.julialang.org/t/batched-matrix-multiply/42332/5 "2020-07-01T07:53:42Z")

</div>

Here are a couple of alternatives. If you use a vector of matrices, that’s faster, otherwise you can get the same effect with `eachslice`, though at some performance cost. And of course, it’s much faster with StaticArrays.

```julia
using BenchmarkTools, Test

A_ = [rand(4,3) for _ in 1:2];
B_ = [rand(3,4) for _ in 1:2];
A = cat(A_...; dims=3)
B = cat(B_...; dims=3)

foo(X, Y) = X .* Y
bar(X, Y) = eachslice(X; dims=3) .* eachslice(Y; dims=3)

```

```julia
julia> @test foo(A_, B_) == bar(A, B)
Test Passed

julia> @btime foo($A_, $B_)
  540.212 ns (3 allocations: 512 bytes)

julia> @btime bar($A, $B);
  1.797 μs (17 allocations: 1.11 KiB)

```

---

<div class="post-metadata">

**Author:** ![RoyiAvital](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/royiavital/32/571_2.png) [@RoyiAvital](https://discourse.julialang.org/u/RoyiAvital)\
**Post date:** [July 1, 2020, 12:26pm UTC](https://discourse.julialang.org/t/batched-matrix-multiply/42332/6 "2020-07-01T12:26:18Z")

</div>

> [@DNF](#):
>
> Here are a couple of alternatives. If you use a vector of matrices, that’s faster, otherwise you can get the same effect with `eachslice`, though at some performance cost. And of course, it’s much faster with StaticArrays.

The idea behind batch multiplication isn’t the coding style of the loop.  
The trick in `MKL` and other libraries implementing Batch Multiplication (Very popular in DL oriented libraries) is getting the computational efficiency of large matrices multiplication. It is done by restructuring the data in a new data form.

---

<div class="post-metadata">

**Author:** ![DNF](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/dnf/32/10191_2.png) [@DNF](https://discourse.julialang.org/u/DNF)\
**Post date:** [July 1, 2020, 12:28pm UTC](https://discourse.julialang.org/t/batched-matrix-multiply/42332/7 "2020-07-01T12:28:30Z")

</div>

I just tried to answer the question in the OP.

---

<div class="post-metadata">

**Author:** ![RoyiAvital](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/royiavital/32/571_2.png) [@RoyiAvital](https://discourse.julialang.org/u/RoyiAvital)\
**Post date:** [July 1, 2020, 1:54pm UTC](https://discourse.julialang.org/t/batched-matrix-multiply/42332/8 "2020-07-01T13:54:54Z")

</div>

Didn’t mean to say you didn’t. I apologize if it was offensive in any way.  
I meant in the context he linked to other discussion where it mentions how batch mode is done correctly by rebuilding the data in a structure which maximizes efficiency of the computation.

---

<div class="post-metadata">

**Author:** ![mcabbott](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/mcabbott/32/6603_2.png) [@mcabbott](https://discourse.julialang.org/u/mcabbott)\
**Post date:** [July 1, 2020, 2:12pm UTC](https://discourse.julialang.org/t/batched-matrix-multiply/42332/9 "2020-07-01T14:12:53Z")

</div>

Has anyone hooked up this MKL batched\_gemm stuff to Julia? If I’m reading correctly, the [link from before](https://software.intel.com/content/www/us/en/develop/articles/introducing-batch-gemm-operations.html) discusses operations which act on an array of matrices, while what you describe sounds more like packed / compact gemm (for many small matrices, stored interleaved).

On the CPU, `batched_mul` is similar to the `eachslice` function above (except that it slices the output too and calls `mul!`). Now that we have [https://github.com/JuliaLang/julia/pull/36360](https://github.com/JuliaLang/julia/pull/36360) it should be upgraded to multi-thread the outer loop.

For tiny matrices like the example above, just writing the loops is faster than calling BLAS in any form. Perhaps StaticArrays would be faster still.

```julia
using NNlib, Einsum
ein(A,B) = @einsum C[i,j,b] := A[i,k,b] * B[k,j,b]
batched_mul(A, B) ≈ ein(A, B) ≈ cat(bar(A,B)...; dims=3)
@code_warntype bar(A, B) # Any

```

```julia
julia> @btime ein($A, $B);
  133.104 ns (1 allocation: 336 bytes)

julia> @btime batched_mul($A, $B);
  328.615 ns (1 allocation: 336 bytes)

```

---

<div class="post-metadata">

**Author:** ![bmit](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/bmit/32/12443_2.png) [@bmit](https://discourse.julialang.org/u/bmit)\
**Post date:** [July 1, 2020, 3:08pm UTC](https://discourse.julialang.org/t/batched-matrix-multiply/42332/10 "2020-07-01T15:08:37Z")

</div>

Thanks for the really fast and thorough answers everyone!

I’m actually using pretty large arrays, so I’d like to avoid copy operations if possible - especially on the GPU. My input is two large N-D arrays, so I’m having trouble converting that to an array of arrays on the GPU without a copy, but I can do it with a `reshape` on an N-D array. There’s probably a trick I’m missing.

At first glance it seems `batched_multiply` might be the best solution for my application, because it works on both CPU and GPU. Ideally, I’d like a function that could also do a broadcast batch multiply. Does anyone know if that exists? Something like

```julia
a = rand(4,3)
b = rand(3,4,2)
batched_multiply(a,b)

```

---

<div class="post-metadata">

**Author:** ![mcabbott](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/mcabbott/32/6603_2.png) [@mcabbott](https://discourse.julialang.org/u/mcabbott)\
**Post date:** [July 1, 2020, 3:22pm UTC](https://discourse.julialang.org/t/batched-matrix-multiply/42332/11 "2020-07-01T15:22:41Z")

</div>

On the GPU, `NNlib.batched_mul` calls `gemm_strided_batched!`, which wants continuous arrays rather than an array of pointers, and is more efficient (IIRC).

This ought to understand broadcasting in the sense of using the same matrix `a` for every slice of 3D `b`, but does not right now. I had a PR to allow this (among other things) [https://github.com/JuliaGPU/CuArrays.jl/pull/664](https://github.com/JuliaGPU/CuArrays.jl/pull/664), after which `batched_mul(reshape(a, size(a)...,1), b)` should work. But I didn’t finish it.

You can of course write `reshape(a * reshape(b, 3,:), 4,4,2)` to do this as one ordinary multiplication. Which (again IIRC) is not as quick for large square-ish arrays as the batched version.

---

<div class="post-metadata">

**Author:** ![RoyiAvital](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/royiavital/32/571_2.png) [@RoyiAvital](https://discourse.julialang.org/u/RoyiAvital)\
**Post date:** [January 31, 2025, 8:51am UTC](https://discourse.julialang.org/t/batched-matrix-multiply/42332/12 "2025-01-31T08:51:46Z")

</div>

> [@mcabbott](#):
>
> Has anyone hooked up this MKL batched\_gemm stuff to Julia? If I’m reading correctly, the [link from before](https://software.intel.com/content/www/us/en/develop/articles/introducing-batch-gemm-operations.html) discusses operations which act on an array of matrices, while what you describe sounds more like packed / compact gemm (for many small matrices, stored interleaved).

Maybe it should be a small grant in the grants program of SciML?

---

<div class="post-metadata">

**Author:** ![draftman9](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/draftman9/32/38128_2.png) [@draftman9](https://discourse.julialang.org/u/draftman9)\
**Post date:** [October 30, 2025, 8:22am UTC](https://discourse.julialang.org/t/batched-matrix-multiply/42332/13 "2025-10-30T08:22:15Z")

</div>

Yes, only for small matrix. For 9x9 matrix, the performance result is reversed.
