# Optimizing Complex Batch Matrix Multiplication

**URL:** <https://discourse.julialang.org/t/optimizing-complex-batch-matrix-multiplication/105381>\
**Category:** Performance\
**Tags:** question\
**Created:** [October 25, 2023, 3:04pm UTC](https://discourse.julialang.org/t/optimizing-complex-batch-matrix-multiplication/105381 "2023-10-25T15:04:35Z")\
**Posts on this page:** 3\
**Page:** 1

<div class="post-metadata">

**Author:** ![quantumtwist](https://avatars.discourse-cdn.com/v4/letter/q/ebca7d/32.png) [@quantumtwist](https://discourse.julialang.org/u/quantumtwist)\
**Post date:** [October 25, 2023, 3:04pm UTC](https://discourse.julialang.org/t/optimizing-complex-batch-matrix-multiplication/105381/1 "2023-10-25T15:04:35Z")

</div>

Hello fellow Julians!  
In my code the bottleneck step is a batched ComplexF64 matrix multiplication of the form `C[i,k,n] = conj(A)[j,i,n] * B[j,k,n].` Naively, one should loop over `n` and do in-place multiplication. However, in the case of floats, it seems that [tullio](https://github.com/mcabbott/Tullio.jl) can speed things up considerably. Is it possible to do something better in the complex case as well? Can you help me optimize things further?

To start off, here’s a reference implementation:

```julia
using LinearAlgebra
using BenchmarkTools
BLAS.set_num_threads(1)
N, M, K = 500, 10, 200
A = rand(ComplexF64,N,M,K)
B = rand(ComplexF64,N,M,K)
C = zeros(ComplexF64,M,M,K);

function batch_zgemm1!(C,A,B)
    for k in axes(C,3)
        @views mul!(C[:,:,k],A[:,:,k]',B[:,:,k])
    end
    return C
end 
@btime batch_zgemm1!($C, $A, $B) # 4.561 ms (0 allocations: 0 bytes)

```

A few notes:

- I require complex matrices with full precision for my use case.
- Typical problem sizes for me are `N = 500-1500`, `M = 2-20`, `K = 100-1600`.
- I am aware that MKL has a batched matrix multiplication library, but it would be good to have a pure-Julia solution for my non-Intel computer.
- The individual multiplications are large enough that BLAS benefits from multithreading, so paralellizing over `K` doesn’t offer any obvious improvements.

---

<div class="post-metadata">

**Author:** ![Oscar\_Smith](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/oscar_smith/32/25343_2.png) [@Oscar\_Smith](https://discourse.julialang.org/u/Oscar_Smith)\
**Post date:** [October 25, 2023, 3:08pm UTC](https://discourse.julialang.org/t/optimizing-complex-batch-matrix-multiplication/105381/2 "2023-10-25T15:08:22Z")

</div>

You want to parellize over k and use single threaded matmul. Matmul paralellizes, but doesn’t do so perfectly.

---

<div class="post-metadata">

**Author:** ![quantumtwist](https://avatars.discourse-cdn.com/v4/letter/q/ebca7d/32.png) [@quantumtwist](https://discourse.julialang.org/u/quantumtwist)\
**Post date:** [October 25, 2023, 3:31pm UTC](https://discourse.julialang.org/t/optimizing-complex-batch-matrix-multiplication/105381/3 "2023-10-25T15:31:16Z")

</div>

I see: you get a speedup of x Nthreads from paralellizing over the outer loop, while BLAS gets an multiplier of \< Nthreads from using more BLAS threads.

```julia
function batch_zgemm2!(C, A, B)
    Threads.@threads for k in axes(A, 3)
        @views mul!(C[:,:,k],A[:,:,k]',B[:,:,k])
    end
    return C
end
Threads.nthreads() # 6
@btime batch_zgemm2!($C, $A, $B) # 758.167 μs (35 allocations: 3.38 KiB)

```

However, is it possible to do better?
