Skip to content

Use the oneMKL trmm/trsm variants for triangular multiply and solve - #643

Merged
michel2323 merged 2 commits into
mainfrom
trmm-trsm-variants
Sep 28, 2026
Merged

michel2323 merged 2 commits into
mainfrom
trmm-trsm-variants

Conversation

@michel2323

Copy link
Copy Markdown
Member

Supersedes #479 (@amontoison's commit is kept as the first commit, rebased onto main with the conflict resolved).

Adds trmm!/trsm!(side, uplo, transa, diag, alpha, [beta,] A, B, C) wrappers for the out-of-place oneMKL variants

C = alpha * op(A) * B + beta * C
C = alpha * op(A) \ B + beta * C

(and the side = 'R' forms) on top of the onemkl{S,D,C,Z}tr{m,s}m_variant entry points already in the support library, and routes generic_trimatmul!, generic_mattrimul!, generic_trimatdiv! and generic_mattridiv! through them. That removes the copyto!(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 aliased B/C would 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 all uplo/trans/diag combinations, beta random and zero, for Float32/Float64/ComplexF32/ComplexF64: all 64 cases correct, B untouched. The onemkl test 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! and gemm are fine. The test suite skips Float64 there, so Buildkite never exercises those types; the ALCF pipeline does.

Changes on top of #479 (second commit)

  • the out-of-place wrappers return C (they returned the untouched B), check size(C) and take the device from A like the in-place routines
  • keep the in-place routines for the aliasing case (see above)
  • drop the mixed triangular x triangular generic_trimatmul! with its throw fallback; the generic path already covers it
  • tests: beta = rand(T) restored for trmm!, right-side trmm! coverage added, and the right-side trsm! 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.jl nor test/onemkl.jl is Runic-clean on main, and the new code follows the surrounding style like the recent changes to these files did.

@github-actions

github-actions Bot commented Sep 24, 2026 •

Copy link
Copy Markdown
Contributor

Your PR requires formatting changes to meet the project's style guidelines.
Please consider running Runic (git runic main) to apply these changes.

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

amontoison and others added 2 commits September 25, 2026 07:47
- 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
@codecov

codecov Bot commented Sep 25, 2026 •

Copy link
Copy Markdown

Codecov Report

✅ All modified and coverable lines are covered by tests.
✅ Project coverage is 80.92%. Comparing base (82be34e) to head (a2c5a6d).

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.
📢 Have feedback on the report? Share it here.

🚀 New features to boost your workflow:
  • ❄️ Test Analytics: Detect flaky tests, report on failures, and find test suite problems.

@michel2323
michel2323 merged commit 6ab8057 into main Sep 28, 2026
5 checks passed
@michel2323
michel2323 deleted the trmm-trsm-variants branch September 28, 2026 13:22
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

None yet

Projects

None yet

Development

Successfully merging this pull request may close these issues.

2 participants