diff --git a/lib/mkl/linalg.jl b/lib/mkl/linalg.jl index 093e3ee7..c65147a7 100644 --- a/lib/mkl/linalg.jl +++ b/lib/mkl/linalg.jl @@ -195,14 +195,22 @@ end end # triangular +# +# `lmul!`/`rmul!`/`ldiv!`/`rdiv!` reach these with the output aliasing the non-triangular +# operand; that case stays on the in-place oneMKL routines. Otherwise the out-of-place +# "variants" write straight into C without first copying the operand over. LinearAlgebra.generic_trimatmul!(C::oneStridedMatrix{T}, uploc, isunitc, tfun::Function, A::oneStridedMatrix{T}, B::oneStridedMatrix{T}) where {T<:onemklFloat} = - trmm!('L', uploc, tfun === identity ? 'N' : tfun === transpose ? 'T' : 'C', isunitc, one(T), A, C === B ? C : copyto!(C, B)) + C === B ? trmm!('L', uploc, tfun === identity ? 'N' : tfun === transpose ? 'T' : 'C', isunitc, one(T), A, C) : + trmm!('L', uploc, tfun === identity ? 'N' : tfun === transpose ? 'T' : 'C', isunitc, one(T), A, B, C) LinearAlgebra.generic_mattrimul!(C::oneStridedMatrix{T}, uploc, isunitc, tfun::Function, A::oneStridedMatrix{T}, B::oneStridedMatrix{T}) where {T<:onemklFloat} = - trmm!('R', uploc, tfun === identity ? 'N' : tfun === transpose ? 'T' : 'C', isunitc, one(T), B, C === A ? C : copyto!(C, A)) + C === A ? trmm!('R', uploc, tfun === identity ? 'N' : tfun === transpose ? 'T' : 'C', isunitc, one(T), B, C) : + trmm!('R', uploc, tfun === identity ? 'N' : tfun === transpose ? 'T' : 'C', isunitc, one(T), B, A, C) LinearAlgebra.generic_trimatdiv!(C::oneStridedMatrix{T}, uploc, isunitc, tfun::Function, A::oneStridedMatrix{T}, B::oneStridedMatrix{T}) where {T<:onemklFloat} = - trsm!('L', uploc, tfun === identity ? 'N' : tfun === transpose ? 'T' : 'C', isunitc, one(T), A, C === B ? C : copyto!(C, B)) + C === B ? trsm!('L', uploc, tfun === identity ? 'N' : tfun === transpose ? 'T' : 'C', isunitc, one(T), A, C) : + trsm!('L', uploc, tfun === identity ? 'N' : tfun === transpose ? 'T' : 'C', isunitc, one(T), A, B, C) LinearAlgebra.generic_mattridiv!(C::oneStridedMatrix{T}, uploc, isunitc, tfun::Function, A::oneStridedMatrix{T}, B::oneStridedMatrix{T}) where {T<:onemklFloat} = - trsm!('R', uploc, tfun === identity ? 'N' : tfun === transpose ? 'T' : 'C', isunitc, one(T), B, C === A ? C : copyto!(C, A)) + C === A ? trsm!('R', uploc, tfun === identity ? 'N' : tfun === transpose ? 'T' : 'C', isunitc, one(T), B, C) : + trsm!('R', uploc, tfun === identity ? 'N' : tfun === transpose ? 'T' : 'C', isunitc, one(T), B, A, C) # # BLAS extensions diff --git a/lib/mkl/wrappers_blas.jl b/lib/mkl/wrappers_blas.jl index 657d177b..afecf7ed 100644 --- a/lib/mkl/wrappers_blas.jl +++ b/lib/mkl/wrappers_blas.jl @@ -1136,6 +1136,80 @@ function trsm(side::Char, trsm!(side, uplo, transa, diag, alpha, A, copy(B)) end +# out-of-place "variants": C = alpha * op(A) * B + beta * C and C = alpha * op(A) \ B + beta * C +# (or B * op(A) resp. B / op(A) for side = 'R'). B is left untouched. +for (mmname_variant, smname_variant, elty) in + ((:onemklDtrmm_variant, :onemklDtrsm_variant, :Float64), + (:onemklStrmm_variant, :onemklStrsm_variant, :Float32), + (:onemklZtrmm_variant, :onemklZtrsm_variant, :ComplexF64), + (:onemklCtrmm_variant, :onemklCtrsm_variant, :ComplexF32)) + @eval begin + function trmm!(side::Char, + uplo::Char, + transa::Char, + diag::Char, + alpha::Number, + beta::Number, + A::oneStridedMatrix{$elty}, + B::oneStridedMatrix{$elty}, + C::oneStridedMatrix{$elty}) + m, n = size(B) + mA, nA = size(A) + if mA != nA throw(DimensionMismatch("A must be square")) end + if nA != (side == 'L' ? m : n) throw(DimensionMismatch("trmm!")) end + if size(C) != (m, n) throw(DimensionMismatch("C must have the same size as B")) end + lda = max(1,stride(A,2)) + ldb = max(1,stride(B,2)) + ldc = max(1,stride(C,2)) + queue = global_queue(context(A), device(A)) + $mmname_variant(sycl_queue(queue), side, uplo, transa, diag, m, n, alpha, A, lda, B, ldb, beta, C, ldc) + C + end + + function trsm!(side::Char, + uplo::Char, + transa::Char, + diag::Char, + alpha::Number, + beta::Number, + A::oneStridedMatrix{$elty}, + B::oneStridedMatrix{$elty}, + C::oneStridedMatrix{$elty}) + m, n = size(B) + mA, nA = size(A) + if mA != nA throw(DimensionMismatch("A must be square")) end + if nA != (side == 'L' ? m : n) throw(DimensionMismatch("trsm!")) end + if size(C) != (m, n) throw(DimensionMismatch("C must have the same size as B")) end + lda = max(1,stride(A,2)) + ldb = max(1,stride(B,2)) + ldc = max(1,stride(C,2)) + queue = global_queue(context(A), device(A)) + $smname_variant(sycl_queue(queue), side, uplo, transa, diag, m, n, alpha, A, lda, B, ldb, beta, C, ldc) + C + end + end +end +function trmm!(side::Char, + uplo::Char, + transa::Char, + diag::Char, + alpha::Number, + A::oneStridedMatrix{T}, + B::oneStridedMatrix{T}, + C::oneStridedMatrix{T}) where T + trmm!(side, uplo, transa, diag, alpha, zero(T), A, B, C) +end +function trsm!(side::Char, + uplo::Char, + transa::Char, + diag::Char, + alpha::Number, + A::oneStridedMatrix{T}, + B::oneStridedMatrix{T}, + C::oneStridedMatrix{T}) where T + trsm!(side, uplo, transa, diag, alpha, zero(T), A, B, C) +end + ## hemm for (fname, elty) in ((:onemklZhemm,:ComplexF64), (:onemklChemm,:ComplexF32)) diff --git a/test/onemkl.jl b/test/onemkl.jl index 253eba4d..10c29521 100644 --- a/test/onemkl.jl +++ b/test/onemkl.jl @@ -673,6 +673,35 @@ end # move to host and compare h_C = Array(dB) @test C ≈ h_C + + dB = oneArray(B) # the in-place call above overwrote dB + C = rand(T,m,n) + dC = oneArray(C) + beta = rand(T) + oneMKL.trmm!('L','U','N','N',alpha,beta,dA,dB,dC) + h_C = Array(dC) + D = alpha*A*B + beta*C + @test D ≈ h_C + @test B == Array(dB) + end + + @testset "right trmm!" begin + A = rand(T,m,m) + B = triu(rand(T, m, m)) + dA = oneArray(A) + dB = oneArray(B) + C = alpha*A*B + dC = copy(dA) + oneMKL.trmm!('R','U','N','N',alpha,dB,dC) + @test C ≈ Array(dC) + + C = rand(T,m,m) + dC = oneArray(C) + beta = rand(T) + oneMKL.trmm!('R','U','T','N',alpha,beta,dB,dA,dC) + h_C = Array(dC) + D = alpha*A*transpose(B) + beta*C + @test D ≈ h_C end @testset "trmm" begin @@ -696,6 +725,15 @@ end dC = copy(dB) oneMKL.trsm!('L','U','N','N',alpha,dA,dC) @test C ≈ Array(dC) + + C = rand(T,m,n) + dC = oneArray(C) + beta = rand(T) + oneMKL.trsm!('L','U','N','N',alpha,beta,dA,dB,dC) + h_C = Array(dC) + D = alpha*(A\B) + beta*C + @test D ≈ h_C + @test B == Array(dB) end @testset "left trsm" begin @@ -737,6 +775,14 @@ end dC = copy(dA) oneMKL.trsm!('R','U','N','N',alpha,dB,dC) @test C ≈ Array(dC) + + C = rand(T,m,m) + dC = oneArray(C) + beta = rand(T) + oneMKL.trsm!('R','U','N','N',alpha,beta,dB,dA,dC) + h_C = Array(dC) + D = alpha*(A/B) + beta*C + @test D ≈ h_C end @testset "right trsm" begin