Skip to content

Commit ebaad02

Browse files
authored
Merge pull request #760 from JuliaGPU/tb/strided_views
Fix multiplication with strided GPU array views
2 parents 1215b4e + 0d956e9 commit ebaad02

2 files changed

Lines changed: 129 additions & 16 deletions

File tree

src/host/linalg.jl

Lines changed: 78 additions & 8 deletions
Original file line numberDiff line numberDiff line change
@@ -467,6 +467,76 @@ end
467467

468468

469469
## matrix multiplication
470+
471+
# GPU-backed members of Base's StridedArray union match LinearAlgebra's BLAS methods,
472+
# but cannot be converted to host pointers. Route them to the generic GPU kernels.
473+
const StridedGPUSubArray{T,N} = Base.StridedSubArray{T,N,<:AbstractGPUArray}
474+
const AnyStridedGPUArray{T,N} = Union{AbstractGPUArray{T,N},StridedGPUSubArray{T,N}}
475+
const AnyStridedGPUVector{T} = AnyStridedGPUArray{T,1}
476+
const AnyStridedGPUMatrix{T} = AnyStridedGPUArray{T,2}
477+
const AnyStridedGPUVecOrMat{T} = Union{AnyStridedGPUVector{T},AnyStridedGPUMatrix{T}}
478+
const AnyStridedGPUMatrixOperand{T} = Union{
479+
AnyStridedGPUMatrix{T},
480+
Adjoint{T,<:AnyStridedGPUMatrix{T}},
481+
Transpose{T,<:AnyStridedGPUMatrix{T}},
482+
Symmetric{T,<:AnyStridedGPUMatrix{T}},
483+
Hermitian{T,<:AnyStridedGPUMatrix{T}},
484+
}
485+
const AnyStridedGPUVecOrMatOperand{T} = Union{
486+
AnyStridedGPUVector{T},
487+
AnyStridedGPUMatrixOperand{T},
488+
}
489+
490+
has_strided_gpu_view(As...) =
491+
any(A -> LinearAlgebra._unwrap(A) isa StridedGPUSubArray, As)
492+
493+
# Intercept before LinearAlgebra unwraps operands and dispatches to backend BLAS methods. The
494+
# signatures also cover view-free products to avoid overlapping methods for each possible view
495+
# position; those calls are sent back through LinearAlgebra's original implementation.
496+
@static if VERSION < v"1.11"
497+
function LinearAlgebra.mul!(C::AnyStridedGPUVector, A::AnyStridedGPUMatrixOperand,
498+
B::AnyStridedGPUVector, a::Number, b::Number)
499+
if has_strided_gpu_view(C, A, B)
500+
return generic_matmatmul!(C, A, B, a, b)
501+
end
502+
invoke(LinearAlgebra.mul!,
503+
Tuple{AbstractVector,LinearAlgebra.AbstractVecOrMat,AbstractVector,Number,Number},
504+
C, A, B, a, b)
505+
end
506+
507+
function LinearAlgebra.mul!(C::AnyStridedGPUMatrix, A::AnyStridedGPUVecOrMatOperand,
508+
B::AnyStridedGPUVecOrMatOperand, a::Number, b::Number)
509+
if has_strided_gpu_view(C, A, B)
510+
return generic_matmatmul!(C, A, B, a, b)
511+
end
512+
invoke(LinearAlgebra.mul!,
513+
Tuple{AbstractMatrix,LinearAlgebra.AbstractVecOrMat,
514+
LinearAlgebra.AbstractVecOrMat,Number,Number},
515+
C, A, B, a, b)
516+
end
517+
else
518+
function LinearAlgebra._mul!(C::AnyStridedGPUVector, A::AnyStridedGPUMatrixOperand,
519+
B::AnyStridedGPUVector, a::Number, b::Number)
520+
if has_strided_gpu_view(C, A, B)
521+
return generic_matmatmul!(C, A, B, a, b)
522+
end
523+
invoke(LinearAlgebra._mul!,
524+
Tuple{AbstractVector,LinearAlgebra.AbstractVecOrMat,AbstractVector,Number,Number},
525+
C, A, B, a, b)
526+
end
527+
528+
function LinearAlgebra._mul!(C::AnyStridedGPUMatrix, A::AnyStridedGPUVecOrMatOperand,
529+
B::AnyStridedGPUVecOrMatOperand, a::Number, b::Number)
530+
if has_strided_gpu_view(C, A, B)
531+
return generic_matmatmul!(C, A, B, a, b)
532+
end
533+
invoke(LinearAlgebra._mul!,
534+
Tuple{AbstractMatrix,LinearAlgebra.AbstractVecOrMat,
535+
LinearAlgebra.AbstractVecOrMat,Number,Number},
536+
C, A, B, a, b)
537+
end
538+
end
539+
470540
# legacy method
471541
generic_matmatmul!(C::AbstractArray, A::AbstractArray, B::AbstractArray, a::Number, b::Number) =
472542
generic_matmatmul!(C, A, B, MulAddMul(a, b))
@@ -500,19 +570,19 @@ function generic_matmatmul!(C::AbstractArray{R}, A::AbstractArray{T}, B::Abstrac
500570
end
501571

502572
@static if !isdefined(LinearAlgebra, Symbol("@stable_muladdmul")) # @stable_muladdmul was added in 1.12
503-
function LinearAlgebra.generic_matvecmul!(C::AbstractGPUVector, tA::AbstractChar, A::AbstractGPUMatrix, B::AbstractGPUVector, _add::MulAddMul = MulAddMul())
573+
function LinearAlgebra.generic_matvecmul!(C::AnyStridedGPUVector, tA::AbstractChar, A::AnyStridedGPUMatrix, B::AnyStridedGPUVector, _add::MulAddMul = MulAddMul())
504574
generic_matmatmul!(C, wrap(A, tA), B, _add)
505575
end
506576

507-
function LinearAlgebra.generic_matmatmul!(C::AbstractGPUVecOrMat, tA, tB, A::AbstractGPUVecOrMat, B::AbstractGPUVecOrMat, _add::MulAddMul=MulAddMul())
577+
function LinearAlgebra.generic_matmatmul!(C::AnyStridedGPUVecOrMat, tA, tB, A::AnyStridedGPUVecOrMat, B::AnyStridedGPUVecOrMat, _add::MulAddMul=MulAddMul())
508578
generic_matmatmul!(C, wrap(A, tA), wrap(B, tB), _add)
509579
end
510580
else
511-
function LinearAlgebra.generic_matvecmul!(C::AbstractGPUVector, tA::AbstractChar, A::AbstractGPUMatrix, B::AbstractGPUVector, a::Number, b::Number)
581+
function LinearAlgebra.generic_matvecmul!(C::AnyStridedGPUVector, tA::AbstractChar, A::AnyStridedGPUMatrix, B::AnyStridedGPUVector, a::Number, b::Number)
512582
LinearAlgebra.@stable_muladdmul generic_matmatmul!(C, wrap(A, tA), B, MulAddMul(a, b))
513583
end
514584

515-
function LinearAlgebra.generic_matmatmul!(C::AbstractGPUVecOrMat, tA, tB, A::AbstractGPUVecOrMat, B::AbstractGPUVecOrMat, a::Number, b::Number)
585+
function LinearAlgebra.generic_matmatmul!(C::AnyStridedGPUVecOrMat, tA, tB, A::AnyStridedGPUVecOrMat, B::AnyStridedGPUVecOrMat, a::Number, b::Number)
516586
LinearAlgebra.@stable_muladdmul generic_matmatmul!(C, wrap(A, tA), wrap(B, tB), MulAddMul(a, b))
517587
end
518588
end
@@ -556,22 +626,22 @@ end
556626
@static if VERSION v"1.12.0-rc"
557627
# we need to use the generic wrapper to avoid dispatch to the 2x2or3x3 method
558628
using LinearAlgebra: generic_matmatmul_wrapper!, BlasFlag
559-
function LinearAlgebra.generic_matmatmul_wrapper!(C::AbstractGPUMatrix{T}, tA::AbstractChar, tB::AbstractChar, A::AbstractGPUVecOrMat{T}, B::AbstractGPUVecOrMat{T}, alpha::Number, beta::Number, val::LinearAlgebra.BlasFlag.SyrkHerkGemm) where {T}
629+
function LinearAlgebra.generic_matmatmul_wrapper!(C::AnyStridedGPUMatrix{T}, tA::AbstractChar, tB::AbstractChar, A::AnyStridedGPUVecOrMat{T}, B::AnyStridedGPUVecOrMat{T}, alpha::Number, beta::Number, val::LinearAlgebra.BlasFlag.SyrkHerkGemm) where {T}
560630
LinearAlgebra.generic_matmatmul!(C, tA, tB, A, B, alpha, beta)
561631
end
562632
# Symmetric/Hermitian inputs with BLAS eltypes would otherwise dispatch to BLAS.symm!/
563633
# hemm!: GPU arrays are DenseArrays, so they match the StridedMatrix{<:BlasFloat} methods
564-
function LinearAlgebra.generic_matmatmul_wrapper!(C::AbstractGPUMatrix{T}, tA::AbstractChar, tB::AbstractChar, A::AbstractGPUVecOrMat{T}, B::AbstractGPUVecOrMat{T}, alpha::Number, beta::Number, val::LinearAlgebra.BlasFlag.SymmHemmGeneric) where {T}
634+
function LinearAlgebra.generic_matmatmul_wrapper!(C::AnyStridedGPUMatrix{T}, tA::AbstractChar, tB::AbstractChar, A::AnyStridedGPUVecOrMat{T}, B::AnyStridedGPUVecOrMat{T}, alpha::Number, beta::Number, val::LinearAlgebra.BlasFlag.SymmHemmGeneric) where {T}
565635
LinearAlgebra.generic_matmatmul!(C, tA, tB, A, B, alpha, beta)
566636
end
567637
# need to support mixed complex/real types too
568638
#function LinearAlgebra.generic_matmatmul_wrapper!(C::AbstractGPUMatrix{Complex{T}}, tA::AbstractChar, tB::AbstractChar, A::AbstractGPUVecOrMat{Complex{T}}, B::AbstractGPUVecOrMat{T}, alpha::Number, beta::Number, val::V) where {T<:BlasReal, V<:LinearAlgebra.BlasFlag.SyrkHerkGemm}
569639
# LinearAlgebra.generic_matmatmul!(C, tA, tB, A, B, alpha, beta)
570640
#end
571-
function LinearAlgebra.generic_matmatmul_wrapper!(C::AbstractGPUMatrix{Complex{T}}, tA::AbstractChar, tB::AbstractChar, A::AbstractGPUVecOrMat{Complex{T}}, B::AbstractGPUVecOrMat{T}, alpha::Number, beta::Number, val::Val{LinearAlgebra.BlasFlag.GEMM}) where T<:Union{Float32, Float64}
641+
function LinearAlgebra.generic_matmatmul_wrapper!(C::AnyStridedGPUMatrix{Complex{T}}, tA::AbstractChar, tB::AbstractChar, A::AnyStridedGPUVecOrMat{Complex{T}}, B::AnyStridedGPUVecOrMat{T}, alpha::Number, beta::Number, val::Val{LinearAlgebra.BlasFlag.GEMM}) where T<:Union{Float32, Float64}
572642
LinearAlgebra.generic_matmatmul!(C, tA, tB, A, B, alpha, beta)
573643
end
574-
function LinearAlgebra.generic_matmatmul_wrapper!(C::AbstractGPUMatrix{Complex{T}}, tA::AbstractChar, tB::AbstractChar, A::AbstractGPUVecOrMat{T}, B::AbstractGPUVecOrMat{Complex{T}}, alpha::Number, beta::Number, val::Val{LinearAlgebra.BlasFlag.GEMM}) where T<:Union{Float32, Float64}
644+
function LinearAlgebra.generic_matmatmul_wrapper!(C::AnyStridedGPUMatrix{Complex{T}}, tA::AbstractChar, tB::AbstractChar, A::AnyStridedGPUVecOrMat{T}, B::AnyStridedGPUVecOrMat{Complex{T}}, alpha::Number, beta::Number, val::Val{LinearAlgebra.BlasFlag.GEMM}) where T<:Union{Float32, Float64}
575645
LinearAlgebra.generic_matmatmul!(C, tA, tB, A, B, alpha, beta)
576646
end
577647
# Julia 1.12 introduced generic_mul! for scalar * array operations

test/testsuite/linalg.jl

Lines changed: 51 additions & 8 deletions
Original file line numberDiff line numberDiff line change
@@ -573,14 +573,57 @@ end
573573
@test compare(mul!, AT, C, A, B, Ref(T(4)), Ref(T(5)))
574574
@test typeof(AT(rand(Tc, 3, 3)) * AT(rand(T, 3, 3))) <: AbstractMatrix
575575
end
576-
@testset "$T with views" for T in eltypes
577-
A = rand(T, 10, 10)
578-
v1 = @view(A[:, 1:5])
579-
v2 = @view(A[1:5, :])
580-
dA = AT(A)
581-
dv1 = @view(dA[:, 1:5])
582-
dv2 = @view(dA[1:5, :])
583-
@test Array(v1) * Array(v2) Array(v1 * v2)
576+
end
577+
578+
@testsuite "linalg/mul!/strided-views" (AT, eltypes)->begin
579+
@testset "$T" for T in (Float16, Float32, ComplexF32)
580+
T in eltypes || continue
581+
582+
A = rand(T, 4, 3)
583+
B = rand(T, 4, 5)
584+
b = rand(T, 4)
585+
586+
@test compare((A, B) -> view(A, 1:2, :)' * view(B, 1:2, :), AT, A, B)
587+
@test compare((A, B) -> transpose(view(A, 1:2, :)) * view(B, 1:2, :), AT, A, B)
588+
@test compare((A, b) -> view(A, 1:2, :)' * view(b, 1:2), AT, A, b)
589+
@test compare(A -> view(A, 1:2, :)' * view(A, 1:2, :), AT, A)
590+
591+
A2 = rand(T, 2, 3)
592+
B2 = rand(T, 2, 5)
593+
b2 = rand(T, 2)
594+
@test compare((A, B) -> view(A, 1:2, :)' * B, AT, A, B2)
595+
@test compare((A, B) -> A' * view(B, 1:2, :), AT, A2, B)
596+
@test compare((A, b) -> view(A, 1:2, :)' * b, AT, A, b2)
597+
@test compare((A, b) -> A' * view(b, 1:2), AT, A2, b)
598+
599+
C = rand(T, 4, 5)
600+
@test compare(AT, C, A, B) do C, A, B
601+
mul!(view(C, 1:3, :), view(A, 1:2, :)', view(B, 1:2, :), T(2), T(1))
602+
C
603+
end
604+
@test compare(AT, C, A2, B2) do C, A, B
605+
mul!(view(C, 1:3, :), A', B, T(2), T(1))
606+
C
607+
end
608+
609+
S = rand(T, 4, 4)
610+
@test compare((S, B) -> Symmetric(view(S, 1:2, 1:2)) * view(B, 1:2, :), AT, S, B)
611+
end
612+
613+
if Float32 in eltypes && ComplexF32 in eltypes
614+
A = rand(ComplexF32, 4, 3)
615+
B = rand(Float32, 3, 6)
616+
@test compare((A, B) -> view(A, 1:2, :) * view(B, :, 1:5), AT, A, B)
617+
618+
A = rand(Float32, 4, 3)
619+
B = rand(ComplexF32, 3, 6)
620+
@test compare((A, B) -> view(A, 1:2, :) * view(B, :, 1:5), AT, A, B)
621+
end
622+
623+
if Int16 in eltypes
624+
A = reshape(Int16.(1:12), 4, 3)
625+
B = reshape(Int16.(1:20), 4, 5)
626+
@test compare((A, B) -> view(A, 1:2, :)' * view(B, 1:2, :), AT, A, B)
584627
end
585628
end
586629

0 commit comments

Comments
 (0)