Skip to content

LinearAlgebra.mul! with Matrix{Dual} does not dispatch to BLAS #854

Description

@abussy

I have noticed that GEMM of matrices involving Dual numbers are not dispatched to BLAS, which significantly impacts performance. This happens with 3- and 5-arguments versions of LinearAlgebra.mul!. Generally, any GEMM of matrices resulting in a Matrix{<:Dual} is impacted, e.g. mul!(C, A, B) will be slow if A isa Matrix{<:Dual} and B 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.

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.")

Activity

Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Metadata

Metadata

Assignees

No one assigned

    Labels

    No labels
    No labels

    Type

    No type

    Projects

    No projects

      Milestone

      No milestone

      Relationships

      None yet

      Development

      No branches or pull requests

      Issue actions