From ae4dc5e30fe41fd1f6f1c8e5876c1628e9e2ce49 Mon Sep 17 00:00:00 2001 From: Alexis Montoison Date: Tue, 2 Sep 2025 11:16:45 -0500 Subject: [PATCH 1/2] [oneMKL] Interface variants of trsm! and trmm! --- lib/mkl/linalg.jl | 8 ++--- lib/mkl/wrappers_blas.jl | 70 ++++++++++++++++++++++++++++++++++++++++ test/onemkl.jl | 24 ++++++++++++++ 3 files changed, 98 insertions(+), 4 deletions(-) diff --git a/lib/mkl/linalg.jl b/lib/mkl/linalg.jl index 093e3ee7..a767325c 100644 --- a/lib/mkl/linalg.jl +++ b/lib/mkl/linalg.jl @@ -196,13 +196,13 @@ end # triangular 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)) + 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)) + 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)) + 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)) + 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..1facf2de 100644 --- a/lib/mkl/wrappers_blas.jl +++ b/lib/mkl/wrappers_blas.jl @@ -1136,6 +1136,76 @@ function trsm(side::Char, trsm!(side, uplo, transa, diag, alpha, A, copy(B)) end +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 + lda = max(1,stride(A,2)) + ldb = max(1,stride(B,2)) + ldc = max(1,stride(C,2)) + queue = global_queue(context(A), device()) + $mmname_variant(sycl_queue(queue), side, uplo, transa, diag, m, n, alpha, A, lda, B, ldb, beta, C, ldc) + B + 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 + lda = max(1,stride(A,2)) + ldb = max(1,stride(B,2)) + ldc = max(1,stride(C,2)) + queue = global_queue(context(A), device()) + $smname_variant(sycl_queue(queue), side, uplo, transa, diag, m, n, alpha, A, lda, B, ldb, beta, C, ldc) + B + 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..57f93c7f 100644 --- a/test/onemkl.jl +++ b/test/onemkl.jl @@ -673,6 +673,14 @@ end # move to host and compare h_C = Array(dB) @test C ≈ h_C + + C = rand(T,m,n) + dC = oneArray(C) + beta = zero(T) # 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 end @testset "trmm" begin @@ -696,6 +704,14 @@ 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 end @testset "left trsm" begin @@ -737,6 +753,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,dA,dB,dC) + h_C = Array(dC) + D = alpha*(A/B) + beta*C + @test D ≈ h_C end @testset "right trsm" begin From a2c5a6dc4343ce0e64b522c846d6fc64f77ea0c3 Mon Sep 17 00:00:00 2001 From: Michel Schanen Date: Thu, 24 Sep 2026 15:30:14 +0000 Subject: [PATCH 2/2] Fix up the trmm/trsm variants - return C from the out-of-place wrappers (they returned the untouched B), check size(C) and take the device from A like the in-place routines - keep the in-place routines when the output aliases the operand (lmul!/rmul!/ldiv!/rdiv! call generic_*! with C === B or C === A); the out-of-place kernel with aliased B/C is a race - drop the mixed triangular x triangular generic_trimatmul! with its throw fallback; the generic path already covers it - tests: beta = rand(T) for trmm!, right trmm! coverage, and the right trsm! variant test passed the triangular matrix in the wrong slot --- lib/mkl/linalg.jl | 16 ++++++++++++---- lib/mkl/wrappers_blas.jl | 12 ++++++++---- test/onemkl.jl | 26 ++++++++++++++++++++++++-- 3 files changed, 44 insertions(+), 10 deletions(-) diff --git a/lib/mkl/linalg.jl b/lib/mkl/linalg.jl index a767325c..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, B, C) + 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, A, C) + 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, B, C) + 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, A, C) + 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 1facf2de..afecf7ed 100644 --- a/lib/mkl/wrappers_blas.jl +++ b/lib/mkl/wrappers_blas.jl @@ -1136,6 +1136,8 @@ 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), @@ -1155,12 +1157,13 @@ for (mmname_variant, smname_variant, elty) in 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()) + 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) - B + C end function trsm!(side::Char, @@ -1176,12 +1179,13 @@ for (mmname_variant, smname_variant, elty) in 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()) + 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) - B + C end end end diff --git a/test/onemkl.jl b/test/onemkl.jl index 57f93c7f..10c29521 100644 --- a/test/onemkl.jl +++ b/test/onemkl.jl @@ -674,13 +674,34 @@ end 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 = zero(T) # rand(T) + 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 @@ -712,6 +733,7 @@ end h_C = Array(dC) D = alpha*(A\B) + beta*C @test D ≈ h_C + @test B == Array(dB) end @testset "left trsm" begin @@ -757,7 +779,7 @@ end C = rand(T,m,m) dC = oneArray(C) beta = rand(T) - oneMKL.trsm!('R','U','N','N',alpha,beta,dA,dB,dC) + 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