# Matrix Multiplication of a Large Number of Small Matrices

**URL:** <https://discourse.julialang.org/t/matrix-multiplication-of-a-large-number-of-small-matrices/93531>\
**Category:** Performance\
**Tags:** linearalgebra, tullio, column-major\
**Created:** [January 25, 2023, 7:45pm UTC](https://discourse.julialang.org/t/matrix-multiplication-of-a-large-number-of-small-matrices/93531 "2023-01-25T19:45:37Z")\
**Posts on this page:** 13\
**Page:** 1

<div class="post-metadata">

**Author:** ![andrewkhardy](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/andrewkhardy/32/27311_2.png) [@andrewkhardy](https://discourse.julialang.org/u/andrewkhardy)\
**Post date:** [January 25, 2023, 7:45pm UTC](https://discourse.julialang.org/t/matrix-multiplication-of-a-large-number-of-small-matrices/93531/1 "2023-01-25T19:45:37Z")

</div>

Im interested in computing the matrix multiplication of 2x2 matrices, but these are sampled at a very large number of points. Here is my current benchmark

```julia
using Tullio,TensorCast, LinearAlgebra, StaticArrays

function MatrixMultiplication(A,B,C)

    @tullio X[k,m,l,i,j] := A[k,i,a] * B[l,k,a,b] * A[k,b,c] *C[m,k,c,j] #- 2*(A[j,k,a] * C[i,j,a,b] * A[j,b,c] *B[m,j,c,l])

    return X;
end;
function RowMatrixMultiplication(A,B,C)
    @tullio X[i,j,k,m,l] := A[i,a,k] * B[a,b,m,k] * A[b,c,k]* C[c,j,l,k] 
    
    return X;
end;

function CastMatrixMultiplication(A,B,C)
    @cast A2[k]{i,j} := A[k,i,j]
    @cast B2[m,k]{i,j} := B[m,k,i,j]
    @cast C2[m,k]{i,j} := C[m,k,i,j]
    @tullio X[m,l,k] := A2[k]*B2[m,k]*A2[k]*C2[l,k]
    return X
end

```

which I call via

```julia
function main()
    A_dim = 6*6
    B_dim = 2^9
    C_dim = B_dim
    A = rand(A_dim,2,2)+rand(A_dim,2,2)*1im
    B = rand(B_dim,A_dim,2,2)+rand(B_dim,A_dim,2,2)*1im
    C = rand(C_dim,A_dim,2,2)+rand(C_dim,A_dim,2,2)*1im
    ws = rand(B_dim)
    @time xa = MatrixMultiplication(A,B,C);
    @time xaa = MatrixMultiplication(A,B,C);
    @time xb = CastMatrixMultiplication(A,B,C);
    @time xbb = CastMatrixMultiplication(A,B,C);
    A = rand(2,2,A_dim)+rand(2,2,A_dim)*1im
    B = rand(2,2,B_dim,A_dim)+rand(2,2,B_dim,A_dim)*1im
    C = rand(2,2,C_dim,A_dim)+rand(2,2,C_dim,A_dim)*1im
    @time xc = RowMatrixMultiplication(A,B,C);
    @time xcc = RowMatrixMultiplication(A,B,C);

 end

```

and I get (for the second run, so without compile time)

```julia
  1.601916 seconds (3 allocations: 576.000 MiB, 9.17% gc time)
  2.227399 seconds (47.19 M allocations: 4.078 GiB, 27.83% gc time)
  1.286178 seconds (3 allocations: 576.000 MiB, 1.95% gc time)

```

My best try in Python is

```julia
def Matrix(A,B,C):
    X = A[None,None,:,:,:] @ B[:,None,:,:,:] @ A[None,None,:,:,:] @ C[None,:,:,:,:]
    return(X)

```

which gives ` ~1.6` seconds as well. The right row first one beats Python somewhat, but I’m used to Tullio seeing 10x increases in speed somewhere. Is this down to BLAS speeds or is there anything I can do? Thanks in advance. I know I can speed it up somewhat by pre-defining the array as well, but am praying for an even more dramatic improvement.

---

<div class="post-metadata">

**Author:** ![Mason](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/mason/32/2423_2.png) [@Mason](https://discourse.julialang.org/u/Mason)\
**Post date:** [January 25, 2023, 8:26pm UTC](https://discourse.julialang.org/t/matrix-multiplication-of-a-large-number-of-small-matrices/93531/2 "2023-01-25T20:26:18Z")

</div>

Julia and python have the opposite major orders, so your `CastMatrixMultiplication` is inefficient. The simplest thing you can do to speed this up substantially is transpose that operation.`

```julia
function CastRowMatrixMultiplication(A,B,C)
    @cast A2[k]{i,j} := A[i,j,k]
    @cast B2[m,k]{i,j} := B[i,j,m,k]
    @cast C2[m,k]{i,j} := C[i,j, m,k]
    @tullio X[m,l,k] := A2[k]*B2[m,k]*A2[k]*C2[l,k]
    return X
end

let
    A_dim = 6*6
    B_dim = 2^9
    C_dim = B_dim
    A = rand(A_dim,2,2)+rand(A_dim,2,2)*1im
    B = rand(B_dim,A_dim,2,2)+rand(B_dim,A_dim,2,2)*1im
    C = rand(C_dim,A_dim,2,2)+rand(C_dim,A_dim,2,2)*1im

    @btime xa = MatrixMultiplication($A,$B,$C);
    @btime xb = CastMatrixMultiplication($A,$B,$C);

    A = rand(2,2,A_dim)+rand(2,2,A_dim)*1im
    B = rand(2,2,B_dim,A_dim)+rand(2,2,B_dim,A_dim)*1im
    C = rand(2,2,C_dim,A_dim)+rand(2,2,C_dim,A_dim)*1im
    
    @btime xcc = RowMatrixMultiplication($A,$B,$C);
    @btime xcc = CastRowMatrixMultiplication($A,$B,$C);
    
 end

```

gives me

```julia
: 325.585 ms (77 allocations: 576.00 MiB) # MatrixMultiplication
: 177.564 ms (95 allocations: 576.01 MiB) # CastMatrixMultiplication
: 226.280 ms (77 allocations: 576.00 MiB) # RowMatrixMultiplication
: 87.582 ms (98 allocations: 576.01 MiB) # CastRowMatrixMultiplication

```

---

<div class="post-metadata">

**Author:** ![andrewkhardy](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/andrewkhardy/32/27311_2.png) [@andrewkhardy](https://discourse.julialang.org/u/andrewkhardy)\
**Post date:** [January 25, 2023, 8:46pm UTC](https://discourse.julialang.org/t/matrix-multiplication-of-a-large-number-of-small-matrices/93531/3 "2023-01-25T20:46:38Z")

</div>

@Mason When I use the @cast versions, I get from using `@btime`

> 2.253 s (47185934 allocations: 4.08 GiB)

for both versions. What version of Julia are you using, or what might I be doing wrong?

---

<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:** [January 25, 2023, 8:50pm UTC](https://discourse.julialang.org/t/matrix-multiplication-of-a-large-number-of-small-matrices/93531/4 "2023-01-25T20:50:54Z")

</div>

> [@andrewkhardy](#):
>
> Im interested in computing the matrix multiplication of 2x2 matrices, but these are sampled at a very large number of points.

[Try StaticArrays](https://docs.julialang.org/en/v1/manual/performance-tips/#Consider-StaticArrays.jl-for-small-fixed-size-vector/matrix-operations).

---

<div class="post-metadata">

**Author:** ![andrewkhardy](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/andrewkhardy/32/27311_2.png) [@andrewkhardy](https://discourse.julialang.org/u/andrewkhardy)\
**Post date:** [January 25, 2023, 8:51pm UTC](https://discourse.julialang.org/t/matrix-multiplication-of-a-large-number-of-small-matrices/93531/5 "2023-01-25T20:51:44Z")

</div>

@stevengj This is exactly what the Cast functions do…

---

<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:** [January 25, 2023, 8:52pm UTC](https://discourse.julialang.org/t/matrix-multiplication-of-a-large-number-of-small-matrices/93531/6 "2023-01-25T20:52:49Z")

</div>

> [@andrewkhardy](#):
>
> This is exactly what the Cast functions do…

To get the full benefit, you need the inputs to be arrays of `SMatrix`, so that the compiler knows the sizes. Don’t convert to static arrays only for the multiplication step, use them from the beginning (and elsewhere throughout your program), e.g. generate them via:

```julia
T = SMatrix{2,2,ComplexF64,4}
A = rand(T,A_dim)
B = rand(T,B_dim,A_dim)
C = rand(T,C_dim,A_dim)

```

This will also fix the issue of the dimension ordering noted by @Mason above (and issues with constant propagation of array sizes, which is no longer needed), though to get [even better locality](https://docs.julialang.org/en/v1/manual/performance-tips/#man-performance-column-major) I would also probably transpose `A_dim` with `B_dim` so that it becomes:

```julia
@tullio X[k,l,m] := A[k]*B[k,m]*A[k]*C[k,l]

```

(but you can experiment a bit with different dimension orderings of `X` and `B` to see which one is fastest with Tullio).

---

<div class="post-metadata">

**Author:** ![Mason](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/mason/32/2423_2.png) [@Mason](https://discourse.julialang.org/u/Mason)\
**Post date:** [January 25, 2023, 9:07pm UTC](https://discourse.julialang.org/t/matrix-multiplication-of-a-large-number-of-small-matrices/93531/7 "2023-01-25T21:07:44Z")

</div>

> [@andrewkhardy](#):
>
> @Mason When I use the @cast versions, I get from using `@btime`
> 
> > 2.253 s (47185934 allocations: 4.08 GiB)
> 
> for both versions. What version of Julia are you using, or what might I be doing wrong?

I’m on 1.9-beta3. You’re probably seeing a type instability because the array sizes aren’t being constant propagated, which if I recall correctly is a fairly recent change.

---

<div class="post-metadata">

**Author:** ![Dan](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/dan/32/42581_2.png) [@Dan](https://discourse.julialang.org/u/Dan)\
**Post date:** [January 25, 2023, 9:08pm UTC](https://discourse.julialang.org/t/matrix-multiplication-of-a-large-number-of-small-matrices/93531/8 "2023-01-25T21:08:28Z")

</div>

I’ve had a good speedup by calculating:

```julia
@tullio AB[k,m] := A[k]*B[k,m]
@tullio AC[k,l] := A[k]*C[k,l]

```

and then doing the cross product (which is much larger):

```julia
@tullio X[k,l,m] := AB[k,m] * AC[k,l]

```

---

<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:** [January 25, 2023, 9:11pm UTC](https://discourse.julialang.org/t/matrix-multiplication-of-a-large-number-of-small-matrices/93531/9 "2023-01-25T21:11:19Z")

</div>

> [@Dan](#):
>
> I’ve had a good speedup by calculating:
> 
> ```julia
> @tullio AB[k,m] := A[k]*B[k,m]
> @tullio AC[k,l] := A[k]*C[k,l]
> 
> ```

Good point, this is a classic space-time tradeoff — it’s a good idea to cache these products with `A` since they are re-used many times. This sort of thing comes up a lot in broadcast-type operations, e.g. [Why is this Julia code considerably slower than Matlab - #62 by stevengj](https://discourse.julialang.org/t/why-is-this-julia-code-considerably-slower-than-matlab/2365/62)

---

<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:** [January 25, 2023, 10:02pm UTC](https://discourse.julialang.org/t/matrix-multiplication-of-a-large-number-of-small-matrices/93531/10 "2023-01-25T22:02:33Z")

</div>

> [@andrewkhardy](#):
>
> `@tullio X[i,j,k,m,l] := A[i,a,k] * B[a,b,m,k] * A[b,c,k]* C[c,j,l,k] `

For large arrays, a chain like this is Tullio’s worst nightmare. It will make 8 nested loops, O(N^8), whereas 3 separate multiplications will be O(N^5) if I’m counting right. For very small arrays, that’s not obviously the right way to think, but I wouldn’t hope for much.

Something similar applies to the arrays of matrices in `A[k]*B[k,m]*A[k]*C[k,l]`, as noted.

> [@Mason](#):
>
> a type instability because the array sizes aren’t being constant propagated,

FWIW I see 2.5s for CastMatrixMultiplication but 160ms for CastRowMatrixMultiplication, on Julia nightly.

TensorCast has a (slightly weird) notation to provide the sizes, if you know they are always 2x2, like so: `@cast B2[m,k]{i:2,j:2} := B[m,k,i,j]`. This removes the instability.

---

<div class="post-metadata">

**Author:** ![andrewkhardy](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/andrewkhardy/32/27311_2.png) [@andrewkhardy](https://discourse.julialang.org/u/andrewkhardy)\
**Post date:** [January 26, 2023, 4:21pm UTC](https://discourse.julialang.org/t/matrix-multiplication-of-a-large-number-of-small-matrices/93531/11 "2023-01-26T16:21:10Z")

</div>

Thanks to @stevengj @Mason, @Dan and @mcabbott, these are all helpful suggestions. I’ve seen a lot of advice about trying to embed as much Julia processing in a single for-loop for optimization reasons. The rest of the process involves

```julia
@reduce XSum[mu,nu] := sum(k,l,m) imag.(X[k,l,m])
result = tr.(XSum)

```

This seems quite tidy, but is there a similar way in that the cross-product gives a ~30% speed-up, I should be unvectorizing some of this for optimization? What’s the intuition (coming from a Python background) about this?

---

<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:** [January 28, 2023, 4:05pm UTC](https://discourse.julialang.org/t/matrix-multiplication-of-a-large-number-of-small-matrices/93531/12 "2023-01-28T16:05:34Z")

</div>

> [@andrewkhardy](#):
>
> ```julia
> XSum[mu,nu] := sum(k,l,m) imag.(X[k,l,m])
> result = tr.(XSum)
> 
> ```

This is pretty wasteful — you are computing the whole Xsum matrix just to take the trace? Also, no need for the dot in the imag call. Why not just sum `tr(imag(X[…]))` directly?

---

<div class="post-metadata">

**Author:** ![andrewkhardy](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/andrewkhardy/32/27311_2.png) [@andrewkhardy](https://discourse.julialang.org/u/andrewkhardy)\
**Post date:** [February 8, 2023, 7:47pm UTC](https://discourse.julialang.org/t/matrix-multiplication-of-a-large-number-of-small-matrices/93531/13 "2023-02-08T19:47:04Z")

</div>

Sorry for bothering folks again, but for memory reasons, I’m attempting to loop over one of these indices in order to reduce the amount of memory.

My comparison is

```julia
using Tullio,TensorCast, LinearAlgebra, StaticArrays, BenchmarkTools

function BestTest(A,B,C)
    @cast A2[k,mu]{i,j} := A[i,j,mu,k]
    @cast B2[k,m]{i,j} := B[i,j,m,k]
    @cast C2[k,l]{i,j} := C[i,j,l,k]
    @tullio AB[k,mu,l] := A2[k,mu] * B2[k,l]
    @tullio AC[k,mu,m] := A2[k,mu] * C2[k,m]
    @tullio X[k,mu,nu,l,m] := AB[k,mu,m] * AC[k,nu,l]
    return X
end
function ForLoopTest(A,B,C)
    X = Array{ComplexF64}(undef,size(A)[4],size(A)[3],size(A)[3])
    #X = Array{SMatrix{2, 2, ComplexF64, 4}}(undef,size(A)[4],size(A)[3],size(A)[3],size(B)[3],size(C)[3])
    for kidx in 1:size(A)[4]
        A1 = A[:,:,:,kidx]
        B1 = B[:,:,:,kidx]
        C1 = C[:,:,:,kidx]
        @cast A2[mu]{i,j}:= A1[i,j,mu]
        @cast B2[m]{i,j} := B1[i,j,m]
        @cast C2[l]{i,j} := C1[i,j,l]
        @tullio AB[mu,l] := A2[mu] * B2[l]
        @tullio AC[mu,m] := A2[mu] * C2[m]
        @tullio X1[mu,nu,l,m] := AB[mu,m] * AC[nu,l]
        @reduce X[kidx,mu,nu] = sum(l,m) tr(X1[mu,nu,l,m] )
    end
    @reduce result[mu,nu] := sum(k) X[k,mu,nu]
    println(result) 
    return result
end

function Combined(A,B,C, Function)
    X = Function(A,B,C)
    @reduce result[mu,nu] := sum(k,l,m) tr(X[k,mu,nu,l,m])
    println(result) 
    return result
end

A_dim = 130
B_dim = 2^8
C_dim = B_dim
A = rand(2,2,2,A_dim)+rand(2,2,2,A_dim)*1im
B = rand(2,2,B_dim,A_dim)+rand(2,2,B_dim,A_dim)*1im
C = rand(2,2,C_dim,A_dim)+rand(2,2,C_dim,A_dim)*1im
@btime x1 = Combined($A,$B,$C, BestTest)
@btime x2 = ForLoopTest($A,$B,$C)

```

Why is the second test 1) 30x slower and 2) seems to give a different answer?
