Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
16 changes: 12 additions & 4 deletions lib/mkl/linalg.jl
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down
74 changes: 74 additions & 0 deletions lib/mkl/wrappers_blas.jl
Original file line number Diff line number Diff line change
Expand Up @@ -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))
Expand Down
46 changes: 46 additions & 0 deletions test/onemkl.jl
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand All @@ -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
Expand Down Expand Up @@ -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
Expand Down
Loading