This is illustrated in the reproducer below for the 3-argument method. To speed things up, matrices are split into values and partials, and all the cross terms explicitly calculated. These are dispatched to BLAS because values and partials are Matrix{<:BlasFloat}. On my machine, I get ~30x speedup for this example. The exact number varies with matrix size and number of partials, but the workaround is systematically faster.
sing ForwardDiff
using LinearAlgebra
# Build random Dual-valued matrices.
function random_dual_matrix(m, n, N)
[ForwardDiff.Dual(rand(), ntuple(_ -> rand(), N)) for _ in 1:m, _ in 1:n]
end
# BLAS-capable baseline for Dual matrices: split into value and each partial,
# call ordinary BLAS-backed mul! on the plain arrays, and reassemble the result.
function mul_dual_blas!(C, A, B)
N = ForwardDiff.npartials(eltype(A))
A_val = ForwardDiff.value.(A)
B_val = ForwardDiff.value.(B)
A_parts = ntuple(i -> ForwardDiff.partials.(A, i), N)
B_parts = ntuple(i -> ForwardDiff.partials.(B, i), N)
C_val = ForwardDiff.value.(C)
C_parts = ntuple(i -> ForwardDiff.partials.(C, i), N)
mul!(C_val, A_val, B_val) # BLAS for values
for i in 1:N
# ∂C_i = ∂A_i * B + A * ∂B_i
mul!(C_parts[i], A_parts[i], B_val)
mul!(C_parts[i], A_val, B_parts[i], true, true)
end
# Reconstruct C from value and partials.
for j in 1:size(C, 2), i in 1:size(C, 1)
C[i, j] = ForwardDiff.Dual(
C_val[i, j],
ForwardDiff.Partials(ntuple(l -> C_parts[l][i, j], N))
)
end
C
end
N_partials = 3
m, n, k = 500, 400, 300
A = random_dual_matrix(m, k, N_partials)
B = random_dual_matrix(k, n, N_partials)
C_generic = similar(A, m, n)
C_blas = similar(A, m, n)
# Warm up both paths before timing.
mul!(C_generic, A, B)
mul!(C_generic, A, B)
mul_dual_blas!(C_blas, A, B)
mul_dual_blas!(C_blas, A, B)
# Time generic 3-arg mul!.
println("Timing generic 3-arg mul! with Matrix{Dual}...")
t_generic = minimum(@elapsed mul!(C_generic, A, B) for _ in 1:5)
# Time BLAS-based baseline.
println("Timing BLAS-based baseline...")
t_blas = minimum(@elapsed mul_dual_blas!(C_blas, A, B) for _ in 1:5)
println()
println("Generic mul! : $(round(t_generic; digits=4)) s")
println("BLAS baseline: $(round(t_blas; digits=4)) s")
println("Speedup : $(round(t_generic / t_blas; digits=2))x")
# Verify correctness.
mul!(C_generic, A, B)
mul_dual_blas!(C_blas, A, B)
@assert C_generic ≈ C_blas
println()
println("Results match.")
I have noticed that GEMM of matrices involving
Dualnumbers are not dispatched to BLAS, which significantly impacts performance. This happens with 3- and 5-arguments versions ofLinearAlgebra.mul!. Generally, any GEMM of matrices resulting in aMatrix{<:Dual}is impacted, e.g.mul!(C, A, B)will be slow ifA isa Matrix{<:Dual}andB isa Matrix{Float64}.This is illustrated in the reproducer below for the 3-argument method. To speed things up, matrices are split into values and partials, and all the cross terms explicitly calculated. These are dispatched to BLAS because values and partials are
Matrix{<:BlasFloat}. On my machine, I get ~30x speedup for this example. The exact number varies with matrix size and number of partials, but the workaround is systematically faster.