Use the oneMKL trmm/trsm variants for triangular multiply and solve - #643
Merged
Merged
Conversation
Contributor
|
Your PR requires formatting changes to meet the project's style guidelines. Click here to view the suggested changes.diff --git a/lib/mkl/linalg.jl b/lib/mkl/linalg.jl
index c65147a..4727441 100644
--- a/lib/mkl/linalg.jl
+++ b/lib/mkl/linalg.jl
@@ -201,16 +201,16 @@ end
# "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} =
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)
+ 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} =
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)
+ 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} =
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)
+ 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} =
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)
+ 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 afecf7e..98e7a22 100644
--- a/lib/mkl/wrappers_blas.jl
+++ b/lib/mkl/wrappers_blas.jl
@@ -1139,75 +1139,97 @@ 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))
+ (
+ (: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})
+ 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))
+ 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
+ return 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})
+ 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))
+ 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
+ return 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)
+function trmm!(
+ side::Char,
+ uplo::Char,
+ transa::Char,
+ diag::Char,
+ alpha::Number,
+ A::oneStridedMatrix{T},
+ B::oneStridedMatrix{T},
+ C::oneStridedMatrix{T}
+ ) where {T}
+ return 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}
+ return trsm!(side, uplo, transa, diag, alpha, zero(T), A, B, C)
end
## hemm
diff --git a/test/onemkl.jl b/test/onemkl.jl
index 10c2952..0a5f5b0 100644
--- a/test/onemkl.jl
+++ b/test/onemkl.jl
@@ -674,34 +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 = 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
+ 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)
+ @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
+ 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
@@ -726,14 +726,14 @@ end
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)
+ 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
@@ -776,13 +776,13 @@ end
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
+ 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 |
- 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
michel2323
force-pushed
the
trmm-trsm-variants
branch
from
September 25, 2026 12:47
6a35bea to
a2c5a6d
Compare
Codecov Report✅ All modified and coverable lines are covered by tests. Additional details and impacted files@@ Coverage Diff @@
## main #643 +/- ##
==========================================
+ Coverage 80.79% 80.92% +0.13%
==========================================
Files 55 55
Lines 4087 4115 +28
==========================================
+ Hits 3302 3330 +28
Misses 785 785 ☔ View full report in Codecov by Harness. 🚀 New features to boost your workflow:
|
This file contains hidden or bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters.
Learn more about bidirectional Unicode characters
Sign up for free
to join this conversation on GitHub.
Already have an account?
Sign in to comment
Add this suggestion to a batch that can be applied as a single commit.This suggestion is invalid because no changes were made to the code.Suggestions cannot be applied while the pull request is closed.Suggestions cannot be applied while viewing a subset of changes.Only one suggestion per line can be applied in a batch.Add this suggestion to a batch that can be applied as a single commit.Applying suggestions on deleted lines is not supported.You must change the existing code in this line in order to create a valid suggestion.Outdated suggestions cannot be applied.This suggestion has been applied or marked resolved.Suggestions cannot be applied from pending reviews.Suggestions cannot be applied on multi-line comments.Suggestions cannot be applied while the pull request is queued to merge.Suggestion cannot be applied right now. Please check back later.
Supersedes #479 (@amontoison's commit is kept as the first commit, rebased onto
mainwith the conflict resolved).Adds
trmm!/trsm!(side, uplo, transa, diag, alpha, [beta,] A, B, C)wrappers for the out-of-place oneMKL variants(and the
side = 'R'forms) on top of theonemkl{S,D,C,Z}tr{m,s}m_variantentry points already in the support library, and routesgeneric_trimatmul!,generic_mattrimul!,generic_trimatdiv!andgeneric_mattridiv!through them. That removes thecopyto!(C, B)that every out-of-place triangular multiply and solve currently pays. When the output aliases the operand (lmul!/rmul!/ldiv!/rdiv!) the in-place routines are kept, since the out-of-place kernel with aliasedB/Cwould be a race.Status of the original blocker
#479 stalled because the variants returned wrong results for
beta != 0(oneMKL 2024.2, 2025.0). With oneMKL 2026.1.0 that is fixed. Checked on an Aurora PVC tile (Max 1550, native fp64, LTS stack): trmm and trsm variants, 4 shapes covering both sides and alluplo/trans/diagcombinations,betarandom and zero, for Float32/Float64/ComplexF32/ComplexF64: all 64 cases correct,Buntouched. Theonemkltest file and a LinearAlgebra dispatch smoke (4 element types x 4 triangular wrappers x 14 operations including the aliasing ones) pass on the same hardware.One caveat: on hardware without native fp64 (Arc A750, fp64 emulation) the Float64/ComplexF64 variants return garbage while the in-place
trmm!/trsm!andgemmare fine. The test suite skips Float64 there, so Buildkite never exercises those types; the ALCF pipeline does.Changes on top of #479 (second commit)
C(they returned the untouchedB), checksize(C)and take the device fromAlike the in-place routinesgeneric_trimatmul!with itsthrowfallback; the generic path already covers itbeta = rand(T)restored fortrmm!, right-sidetrmm!coverage added, and the right-sidetrsm!variant test passed the triangular matrix in the wrong slot (likely the cause of the CI failures in the last round of [oneMKL] Interface variants of trsm! and trmm! #479)Not Runic-formatted on purpose: neither
lib/mkl/wrappers_blas.jlnortest/onemkl.jlis Runic-clean onmain, and the new code follows the surrounding style like the recent changes to these files did.