diff --git a/src/linalg.jl b/src/linalg.jl index d3bf24f3..9b584990 100644 --- a/src/linalg.jl +++ b/src/linalg.jl @@ -359,6 +359,11 @@ end *(A::SparseOrTri, B::AdjOrTrans{<:Any,<:AbstractSparseMatrixCSC}) = spmatmul(A, copy(B)) *(A::AdjOrTrans{<:Any,<:AbstractSparseMatrixCSC}, B::SparseOrTri) = spmatmul(copy(A), B) *(A::AdjOrTrans{<:Any,<:AbstractSparseMatrixCSC}, B::AdjOrTrans{<:Any,<:AbstractSparseMatrixCSC}) = spmatmul(copy(A), copy(B)) +# a symmetric/Hermitian sparse factor is materialized, which is O(nnz) like the copies above +*(A::SparseMatrixCSCSymmHerm, B::Union{SparseOrTri,AdjOrTrans{<:Any,<:AbstractSparseMatrixCSC}}) = sparse(A) * B +*(A::Union{SparseOrTri,AdjOrTrans{<:Any,<:AbstractSparseMatrixCSC}}, B::SparseMatrixCSCSymmHerm) = A * sparse(B) +*(A::SparseMatrixCSCSymmHerm, B::SparseMatrixCSCSymmHerm) = sparse(A) * sparse(B) +*(A::SparseMatrixCSCSymmHerm, x::SparseVectorOrView) = sparse(A) * x (*)(Da::Diagonal, A::Union{SparseMatrixCSCOrView, AdjOrTrans{<:Any,<:AbstractSparseMatrixCSC}}, Db::Diagonal) = Da * (A * Db) function (*)(Da::Diagonal, A::SparseMatrixCSC, Db::Diagonal) diff --git a/test/linalg_products.jl b/test/linalg_products.jl index 20103eae..76227f2e 100644 --- a/test/linalg_products.jl +++ b/test/linalg_products.jl @@ -93,6 +93,30 @@ end end end +@testset "symmetric/Hermitian sparse times sparse" begin + n = 10 + @testset "$T" for T in (Float64, ComplexF64) + A = sprandn(T, n, n, 0.3); B = sprandn(T, n, n, 0.3) + for S in (Symmetric(A), Hermitian(A, :L), Symmetric(view(A, :, 1:n))) + for X in (B, B', transpose(B), UpperTriangular(B), view(B, :, 1:n), Hermitian(B), Symmetric(B, :L)) + @test (S * X)::SparseMatrixCSC ≈ Matrix(S) * Matrix(X) + @test (X * S)::SparseMatrixCSC ≈ Matrix(X) * Matrix(S) + end + for x in (sprandn(T, n, 0.5), view(B, :, 2)) + @test (S * x)::SparseVector ≈ Matrix(S) * Vector(x) + end + end + end + # one multiplication per pair of matching stored entries, not per element + P = mulcount_sparse(sparse(1.0I, n, n)) + for f in (() -> Symmetric(P) * P, () -> P * Symmetric(P), () -> P' * Symmetric(P), + () -> UpperTriangular(P) * Symmetric(P), () -> Symmetric(P) * Symmetric(P, :L)) + @test mulcount(f) == n + end + x = sparsevec(fill(MulCount(1.0), n)) + @test mulcount(() -> Symmetric(P) * x) == mulcount(() -> P * x) +end + @testset "Adding sparse-backed SymTridiagonal (#46355)" begin a = SymTridiagonal(sparsevec(Int[1]), sparsevec(Int[])) @test a + a == Matrix(a) + Matrix(a)