Skip to content

Commit 83c9130

Browse files
authored
Implement LinearAlgebra's storage-specific mul! methods (#630)
1 parent db8bc78 commit 83c9130

3 files changed

Lines changed: 44 additions & 26 deletions

File tree

Project.toml

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -38,7 +38,7 @@ AcceleratedKernels = "0.3.1, 0.4"
3838
Adapt = "4"
3939
CEnum = "0.4, 0.5"
4040
ExprTools = "0.1"
41-
GPUArrays = "11.2.1"
41+
GPUArrays = "11.5.14"
4242
GPUCompiler = "2"
4343
GPUToolbox = "0.1, 0.2, 0.3, 1, 3"
4444
KernelAbstractions = "0.9.39"

lib/mkl/interfaces.jl

Lines changed: 23 additions & 14 deletions
Original file line numberDiff line numberDiff line change
@@ -2,40 +2,49 @@
22

33
using LinearAlgebra: BlasComplex, BlasFloat, BlasReal, MulAddMul
44

5-
# legacy methods with final MulAddMul argument
6-
LinearAlgebra.generic_matvecmul!(C::oneVector{T}, tA::AbstractChar, A::oneSparseMatrixCSR{T}, B::oneVector{T}, _add::MulAddMul) where {T <: Union{Float16, ComplexF16, BlasFloat}} =
7-
LinearAlgebra.generic_matvecmul!(C, tA, A, B, _add.alpha, _add.beta)
8-
LinearAlgebra.generic_matvecmul!(C::oneVector{T}, tA::AbstractChar, A::oneSparseMatrixCSC{T}, B::oneVector{T}, _add::MulAddMul) where {T <: Union{Float16, ComplexF16, BlasFloat}} =
9-
LinearAlgebra.generic_matvecmul!(C, tA, A, B, _add.alpha, _add.beta)
10-
LinearAlgebra.generic_matmatmul!(C::oneMatrix{T}, tA, tB, A::oneSparseMatrixCSR{T}, B::oneMatrix{T}, _add::MulAddMul) where {T <: Union{Float16, ComplexF16, BlasFloat}} =
11-
LinearAlgebra.generic_matmatmul!(C, tA, tB, A, B, _add.alpha, _add.beta)
12-
LinearAlgebra.generic_matmatmul!(C::oneMatrix{T}, tA, tB, A::oneSparseMatrixCSC{T}, B::oneMatrix{T}, _add::MulAddMul) where {T <: Union{Float16, ComplexF16, BlasFloat}} =
13-
LinearAlgebra.generic_matmatmul!(C, tA, tB, A, B, _add.alpha, _add.beta)
14-
15-
function LinearAlgebra.generic_matvecmul!(C::oneVector{T}, tA::AbstractChar, A::oneSparseMatrixCSR{T}, B::oneVector{T}, alpha::Number, beta::Number) where {T <: BlasFloat}
5+
function LinearAlgebra.mul!(C::oneVector{T}, tA::AbstractChar, A::oneSparseMatrixCSR{T}, B::oneVector{T}, alpha::Number, beta::Number) where {T <: BlasFloat}
166
tA = tA in ('S', 's', 'H', 'h') ? 'N' : tA
177
return sparse_gemv!(tA, alpha, A, B, beta, C)
188
end
199

20-
function LinearAlgebra.generic_matvecmul!(C::oneVector{T}, tA::AbstractChar, A::oneSparseMatrixCSC{T}, B::oneVector{T}, alpha::Number, beta::Number) where {T <: BlasFloat}
10+
function LinearAlgebra.mul!(C::oneVector{T}, tA::AbstractChar, A::oneSparseMatrixCSC{T}, B::oneVector{T}, alpha::Number, beta::Number) where {T <: BlasFloat}
2111
# sparse_gemv! already maps op(A) onto the transposed CSR handle, so tA is passed through
2212
tA = tA in ('S', 's', 'H', 'h') ? 'N' : tA
2313
return sparse_gemv!(tA, alpha, A, B, beta, C)
2414
end
2515

26-
function LinearAlgebra.generic_matmatmul!(C::oneMatrix{T}, tA, tB, A::oneSparseMatrixCSR{T}, B::oneMatrix{T}, alpha::Number, beta::Number) where {T <: BlasFloat}
16+
function LinearAlgebra.mul!(C::oneMatrix{T}, tA, tB, A::oneSparseMatrixCSR{T}, B::oneMatrix{T}, alpha::Number, beta::Number) where {T <: BlasFloat}
2717
tA = tA in ('S', 's', 'H', 'h') ? 'N' : tA
2818
tB = tB in ('S', 's', 'H', 'h') ? 'N' : tB
2919
return sparse_gemm!(tA, tB, alpha, A, B, beta, C)
3020
end
3121

32-
function LinearAlgebra.generic_matmatmul!(C::oneMatrix{T}, tA, tB, A::oneSparseMatrixCSC{T}, B::oneMatrix{T}, alpha::Number, beta::Number) where {T <: BlasFloat}
22+
function LinearAlgebra.mul!(C::oneMatrix{T}, tA, tB, A::oneSparseMatrixCSC{T}, B::oneMatrix{T}, alpha::Number, beta::Number) where {T <: BlasFloat}
3323
# sparse_gemm! already maps op(A) onto the transposed CSR handle, so tA is passed through
3424
tA = tA in ('S', 's', 'H', 'h') ? 'N' : tA
3525
tB = tB in ('S', 's', 'H', 'h') ? 'N' : tB
3626
return sparse_gemm!(tA, tB, alpha, A, B, beta, C)
3727
end
3828

29+
# Julia < 1.13 dispatches on the non-public `generic_matvecmul!` and `generic_matmatmul!`,
30+
# which JuliaLang/LinearAlgebra.jl#1671 superseded by the `mul!` methods above. Forward from
31+
# the old names, both the alpha/beta variants (1.12) and the ones taking a final MulAddMul
32+
# (1.10 and 1.11).
33+
@static if VERSION < v"1.13.0-rc4"
34+
for SparseMatrixType in (:oneSparseMatrixCSR, :oneSparseMatrixCSC)
35+
@eval begin
36+
LinearAlgebra.generic_matvecmul!(C::oneVector{T}, tA::AbstractChar, A::$SparseMatrixType{T}, B::oneVector{T}, alpha::Number, beta::Number) where {T <: BlasFloat} =
37+
LinearAlgebra.mul!(C, tA, A, B, alpha, beta)
38+
LinearAlgebra.generic_matvecmul!(C::oneVector{T}, tA::AbstractChar, A::$SparseMatrixType{T}, B::oneVector{T}, _add::MulAddMul) where {T <: BlasFloat} =
39+
LinearAlgebra.mul!(C, tA, A, B, _add.alpha, _add.beta)
40+
LinearAlgebra.generic_matmatmul!(C::oneMatrix{T}, tA, tB, A::$SparseMatrixType{T}, B::oneMatrix{T}, alpha::Number, beta::Number) where {T <: BlasFloat} =
41+
LinearAlgebra.mul!(C, tA, tB, A, B, alpha, beta)
42+
LinearAlgebra.generic_matmatmul!(C::oneMatrix{T}, tA, tB, A::$SparseMatrixType{T}, B::oneMatrix{T}, _add::MulAddMul) where {T <: BlasFloat} =
43+
LinearAlgebra.mul!(C, tA, tB, A, B, _add.alpha, _add.beta)
44+
end
45+
end
46+
end
47+
3948
function LinearAlgebra.generic_trimatdiv!(C::oneVector{T}, uploc, isunitc, tfun::Function, A::oneSparseMatrixCSR{T}, B::oneVector{T}) where {T <: BlasFloat}
4049
return sparse_trsv!(uploc, tfun === identity ? 'N' : tfun === transpose ? 'T' : 'C', isunitc, one(T), A, B, C)
4150
end

lib/mkl/linalg.jl

Lines changed: 20 additions & 11 deletions
Original file line numberDiff line numberDiff line change
@@ -70,9 +70,7 @@ end
7070
#
7171
# BLAS 2
7272

73-
LinearAlgebra.generic_matvecmul!(Y::oneVector, tA::AbstractChar, A::oneStridedMatrix, B::oneStridedVector, _add::MulAddMul) =
74-
LinearAlgebra.generic_matvecmul!(Y, tA, A, B, _add.alpha, _add.beta)
75-
function LinearAlgebra.generic_matvecmul!(Y::oneVector, tA::AbstractChar, A::oneStridedMatrix, B::oneStridedVector, a::Number, b::Number)
73+
function LinearAlgebra.mul!(Y::oneVector, tA::AbstractChar, A::oneStridedMatrix, B::oneStridedVector, a::Number, b::Number)
7674
mA, nA = tA == 'N' ? size(A) : reverse(size(A))
7775

7876
if nA != length(B)
@@ -96,8 +94,8 @@ function LinearAlgebra.generic_matvecmul!(Y::oneVector, tA::AbstractChar, A::one
9694
if tA in ('N', 'T', 'C')
9795
return gemv!(tA, alpha, A, B, beta, Y)
9896
elseif tA in ('S', 's') && T <: Real
99-
# complex symv! is not wrapped; fall through to generic_matmatmul!,
100-
# which can use symm! instead
97+
# complex symv! is not wrapped; fall through to the matrix-matrix
98+
# `mul!`, which can use symm! instead
10199
return symv!(tA == 'S' ? 'U' : 'L', alpha, A, B, beta, Y)
102100
elseif tA in ('H', 'h')
103101
# hemv! only supports complex eltypes, but a real Hermitian matrix
@@ -107,7 +105,7 @@ function LinearAlgebra.generic_matvecmul!(Y::oneVector, tA::AbstractChar, A::one
107105
end
108106
end
109107
end
110-
return LinearAlgebra.generic_matmatmul!(Y, tA, 'N', A, B, alpha, beta)
108+
return LinearAlgebra.mul!(Y, tA, 'N', A, B, alpha, beta)
111109
end
112110

113111
# triangular
@@ -123,11 +121,7 @@ LinearAlgebra.generic_trimatdiv!(C::oneStridedVector{T}, uploc, isunitc, tfun::F
123121
# BLAS 3
124122
#
125123

126-
LinearAlgebra.generic_matmatmul!(
127-
C::oneStridedVecOrMat, tA, tB, A::oneStridedVecOrMat,
128-
B::oneStridedVecOrMat, _add::MulAddMul,
129-
) = LinearAlgebra.generic_matmatmul!(C, tA, tB, A, B, _add.alpha, _add.beta)
130-
function LinearAlgebra.generic_matmatmul!(
124+
function LinearAlgebra.mul!(
131125
C::oneStridedVecOrMat, tA, tB, A::oneStridedVecOrMat,
132126
B::oneStridedVecOrMat, alpha::Number, beta::Number,
133127
)
@@ -185,6 +179,21 @@ function LinearAlgebra.generic_matmatmul!(
185179
GPUArrays.generic_matmatmul!(C, wrap(A, tA), wrap(B, tB), alpha, beta)
186180
end
187181

182+
# Julia < 1.13 dispatches on the non-public `generic_matvecmul!` and `generic_matmatmul!`,
183+
# which JuliaLang/LinearAlgebra.jl#1671 superseded by the `mul!` methods above. Forward from
184+
# the old names, both the alpha/beta variants (1.12) and the ones taking a final MulAddMul
185+
# (1.10 and 1.11).
186+
@static if VERSION < v"1.13.0-rc4"
187+
LinearAlgebra.generic_matvecmul!(Y::oneVector, tA::AbstractChar, A::oneStridedMatrix, B::oneStridedVector, alpha::Number, beta::Number) =
188+
LinearAlgebra.mul!(Y, tA, A, B, alpha, beta)
189+
LinearAlgebra.generic_matvecmul!(Y::oneVector, tA::AbstractChar, A::oneStridedMatrix, B::oneStridedVector, _add::MulAddMul) =
190+
LinearAlgebra.mul!(Y, tA, A, B, _add.alpha, _add.beta)
191+
LinearAlgebra.generic_matmatmul!(C::oneStridedVecOrMat, tA, tB, A::oneStridedVecOrMat, B::oneStridedVecOrMat, alpha::Number, beta::Number) =
192+
LinearAlgebra.mul!(C, tA, tB, A, B, alpha, beta)
193+
LinearAlgebra.generic_matmatmul!(C::oneStridedVecOrMat, tA, tB, A::oneStridedVecOrMat, B::oneStridedVecOrMat, _add::MulAddMul) =
194+
LinearAlgebra.mul!(C, tA, tB, A, B, _add.alpha, _add.beta)
195+
end
196+
188197
# triangular
189198
LinearAlgebra.generic_trimatmul!(C::oneStridedMatrix{T}, uploc, isunitc, tfun::Function, A::oneStridedMatrix{T}, B::oneStridedMatrix{T}) where {T<:onemklFloat} =
190199
trmm!('L', uploc, tfun === identity ? 'N' : tfun === transpose ? 'T' : 'C', isunitc, one(T), A, C === B ? C : copyto!(C, B))

0 commit comments

Comments
 (0)