Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
Show all changes
46 commits
Select commit Hold shift + click to select a range
7475023
Batched SVD support for ROCSOLVER and CUSOLVER
Aug 21, 2026
606c42b
Go back to QRIteration default
Aug 22, 2026
57b5bc9
Fix default driver for CUDA
kshyatt Aug 23, 2026
7ff48b3
Fix typo
kshyatt Aug 23, 2026
e626f2b
And fix bad check in yacusolver
kshyatt Aug 23, 2026
b681779
One more API fix
kshyatt Aug 23, 2026
4e59aaf
Fix size of U and V bar
kshyatt Aug 23, 2026
f8f46f2
More API fixes
kshyatt Aug 23, 2026
e0f4f8b
Remove redundant dummy one liners
kshyatt Sep 14, 2026
a02ce63
Switch to batched_svd_f and get rid of Batched versions of algorithms
kshyatt Sep 14, 2026
e8834e7
Fixup AMD default algos
kshyatt Sep 14, 2026
721c09e
Move ragged batch handling and cutoffs etc into MAK
kshyatt Sep 15, 2026
dbedd55
Add support for a batched svd_via_adjoint for algos requiring tall ma…
Sep 15, 2026
8e4bcc4
Add batched svd_full too
kshyatt Sep 15, 2026
dc48639
Add special path for AMDGPU, implement gesvdx batching there, and fix…
Sep 16, 2026
da3adb8
Some fixes and actually test Bisection on AMD
Sep 16, 2026
d3d4fd4
A few more fixes for Bisection
Sep 17, 2026
91b3dc8
Support svd_full for Bisection
Sep 17, 2026
a2a45f0
Remove batched and non-batched bisection for now
kshyatt Sep 18, 2026
97c4b7b
Restore Bisection test
kshyatt Sep 24, 2026
1db311e
Avoid some allocs
kshyatt Sep 24, 2026
379538b
one to one-liner and more checks for batch_size
kshyatt Sep 25, 2026
4f4974e
Apply batched suggestions from code review
kshyatt Sep 28, 2026
7a77866
Fix sizes of dummy arrays
kshyatt Sep 28, 2026
fcb5c9c
Don't return foreach
kshyatt Sep 28, 2026
ead9a69
Missing batched Bisection piping
Sep 29, 2026
0995349
Refactor out gaugefixing for batches
Sep 29, 2026
bfe9245
Add a flag for whether the driver supports ragged batches
kshyatt Sep 29, 2026
4d02a3c
Dumb typos
kshyatt Sep 29, 2026
1e9646e
Have batched_svd_compact use Diagonal
kshyatt Sep 29, 2026
25529da
Fix dimension check in gesvdx_strided_batched!
Sep 29, 2026
23f4870
supports_ragged_batch takes a driver, not an algorithm
Sep 29, 2026
04df157
Make sure to mark bisection as supporting svd full here
Sep 29, 2026
b9ddc59
Move some batching utilities into a common file
kshyatt Sep 29, 2026
d404f67
Move supports_ragged_batch to Driver definition
kshyatt Sep 29, 2026
d9d8c5e
Just put batches below algos
kshyatt Sep 29, 2026
dfada5b
Put the basis completion in the implementation file
Sep 30, 2026
bb619b5
Move batch_size length checks into else
kshyatt Sep 30, 2026
3aef615
Apply batched suggestions from code review
kshyatt Sep 30, 2026
0b7a6cc
Strides for safety
kshyatt Sep 30, 2026
7cbfec5
Use both alg and driver in the batch checks
kshyatt Sep 30, 2026
382f50e
batch_size checks for batched gesvdx
Oct 1, 2026
818792a
Force driver lookup for support checks
kshyatt Oct 1, 2026
27df039
Dumb typo
kshyatt Oct 1, 2026
e37c010
Actually make sure matrices that are too big aren't sent to the batch…
kshyatt Oct 1, 2026
bb0d96d
Move batch_size check for CUSOLVER too
kshyatt Oct 1, 2026
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
56 changes: 48 additions & 8 deletions ext/MatrixAlgebraKitAMDGPUExt/MatrixAlgebraKitAMDGPUExt.jl
Original file line number Diff line number Diff line change
Expand Up @@ -6,18 +6,26 @@ using MatrixAlgebraKit: one!, zero!, uppertriangular!, lowertriangular!
using MatrixAlgebraKit: diagview, sign_safe
using MatrixAlgebraKit: ROCSOLVER, LQViaTransposedQR, TruncationStrategy, NoTruncation, TruncationByValue, AbstractAlgorithm
using MatrixAlgebraKit: default_qr_algorithm, default_lq_algorithm, default_svd_algorithm, default_eigh_algorithm
import MatrixAlgebraKit: geqrf!, ungqr!, unmqr!, gesvd!, gesvdx!, gesvdj!
import MatrixAlgebraKit: geqrf!, ungqr!, unmqr!, gesvd!, gesdd!, gesvdx!, gesvdj!
import MatrixAlgebraKit: gesvdj_batched!, gesdd_batched!, gesvd_batched!, gesvdx_batched!
import MatrixAlgebraKit: heevj!, heevd!, heev!, heevx!
import MatrixAlgebraKit: _sylvester, svd_rank, svd_pullback!, complete_svd_basis!
import MatrixAlgebraKit: _sylvester, svd_rank, svd_pullback!
using AMDGPU
using LinearAlgebra
using LinearAlgebra: BlasFloat

include("yarocsolver.jl")

MatrixAlgebraKit.default_driver(::Type{TA}) where {TA <: StridedROCVecOrMat{<:BlasFloat}} = ROCSOLVER()
MatrixAlgebraKit.default_driver(::Type{TA}) where {TA <: StridedROCArray{<:BlasFloat}} = ROCSOLVER()
MatrixAlgebraKit.default_driver(::Type{TA}) where {TA <: AbstractVector{<:StridedROCMatrix{<:BlasFloat}}} = ROCSOLVER()

function MatrixAlgebraKit.default_svd_algorithm(::Type{T}; kwargs...) where {T <: StridedROCVecOrMat{<:BlasFloat}}
function MatrixAlgebraKit.default_svd_algorithm(::Type{T}; kwargs...) where {T <: StridedROCMatrix{<:BlasFloat}}
return QRIteration(; kwargs...)
end
function MatrixAlgebraKit.default_svd_algorithm(::Type{T}; kwargs...) where {T <: StridedROCArray{<:BlasFloat, 3}}
return QRIteration(; kwargs...)
end
function MatrixAlgebraKit.default_svd_algorithm(::Type{T}; kwargs...) where {T <: AbstractVector{<:StridedROCMatrix{<:BlasFloat}}}
return QRIteration(; kwargs...)
end
function MatrixAlgebraKit.default_eigh_algorithm(::Type{T}; kwargs...) where {T <: StridedROCVecOrMat{<:BlasFloat}}
Expand All @@ -28,7 +36,12 @@ for f in (:geqrf!, :ungqr!, :unmqr!)
@eval $f(::ROCSOLVER, args...) = YArocSOLVER.$f(args...)
end

MatrixAlgebraKit.supports_svd_full(::ROCSOLVER, f::Symbol) = f in (:qr_iteration, :jacobi, :bisection)
MatrixAlgebraKit.supports_svd_full(::ROCSOLVER, f::Symbol) = f in (:qr_iteration, :jacobi, :divide_and_conquer, :bisection)

# rocSOLVER's `gesvd*_batched` functions take the batch as an array
# of device pointers, so a group of equally sized matrices
# doesn't need a copy into a 3D ROCArray.
MatrixAlgebraKit.supports_pointer_batch(::AbstractAlgorithm, ::ROCSOLVER, ::Type{<:StridedROCMatrix{<:BlasFloat}}) = true

function gesvd!(::ROCSOLVER, A::StridedROCMatrix, S::StridedROCVector, U::StridedROCMatrix, Vᴴ::StridedROCMatrix; kwargs...)
m, n = size(A)
Expand All @@ -42,11 +55,38 @@ function gesvdj!(::ROCSOLVER, A::StridedROCMatrix, S::StridedROCVector, U::Strid
return MatrixAlgebraKit.svd_via_adjoint!(gesvdj!, ROCSOLVER(), A, S, U, Vᴴ; kwargs...)
end

function gesvdx!(::ROCSOLVER, A::StridedROCMatrix, S::StridedROCVector, U::StridedROCMatrix, Vᴴ::StridedROCMatrix; kwargs...)
gesvdx!(::ROCSOLVER, A::StridedROCMatrix, S::StridedROCVector, U::StridedROCMatrix, Vᴴ::StridedROCMatrix; kwargs...) =
YArocSOLVER.gesvdx!(A, S, U, Vᴴ; kwargs...)
complete_svd_basis!(U, Vᴴ, length(S))
return S, U, Vᴴ

# rocSOLVER's batched `gesvd` requires m ≥ n, so wide matrices go through the adjoint
function gesvd_batched!(::ROCSOLVER, As::AbstractVector{<:StridedROCMatrix}, Ss::StridedROCMatrix, Us::StridedROCArray{T, 3}, Vᴴs::StridedROCArray{T, 3}; kwargs...) where {T <: BlasFloat}
Comment thread
kshyatt marked this conversation as resolved.
m, n = size(first(As))
m >= n && return YArocSOLVER.gesvd_batched!(As, Ss, Us, Vᴴs; kwargs...)
return MatrixAlgebraKit.batched_svd_via_adjoint!(gesvd_batched!, ROCSOLVER(), As, Ss, Us, Vᴴs; kwargs...)
end
function gesvd_batched!(::ROCSOLVER, As::StridedROCArray{T, 3}, Ss::StridedROCMatrix, Us::StridedROCArray{T, 3}, Vᴴs::StridedROCArray{T, 3}; kwargs...) where {T <: BlasFloat}
m, n, _ = size(As)
m >= n && return YArocSOLVER.gesvd_strided_batched!(As, Ss, Us, Vᴴs; kwargs...)
return MatrixAlgebraKit.batched_svd_via_adjoint!(gesvd_batched!, ROCSOLVER(), As, Ss, Us, Vᴴs; kwargs...)
end

gesdd_batched!(::ROCSOLVER, As::AbstractVector{<:StridedROCMatrix}, Ss::StridedROCMatrix, Us::StridedROCArray{T, 3}, Vᴴs::StridedROCArray{T, 3}; kwargs...) where {T <: BlasFloat} =
YArocSOLVER.gesdd_batched!(As, Ss, Us, Vᴴs; kwargs...)
gesdd_batched!(::ROCSOLVER, As::StridedROCArray{T, 3}, Ss::StridedROCMatrix, Us::StridedROCArray{T, 3}, Vᴴs::StridedROCArray{T, 3}; kwargs...) where {T <: BlasFloat} =
YArocSOLVER.gesdd_strided_batched!(As, Ss, Us, Vᴴs; kwargs...)

gesvdj_batched!(::ROCSOLVER, As::AbstractVector{<:StridedROCMatrix}, Ss::StridedROCMatrix, Us::StridedROCArray{T, 3}, Vᴴs::StridedROCArray{T, 3}; kwargs...) where {T <: BlasFloat} =
YArocSOLVER.gesvdj_batched!(As, Ss, Us, Vᴴs; kwargs...)
gesvdj_batched!(::ROCSOLVER, As::StridedROCArray{T, 3}, Ss::StridedROCMatrix, Us::StridedROCArray{T, 3}, Vᴴs::StridedROCArray{T, 3}; kwargs...) where {T <: BlasFloat} =
YArocSOLVER.gesvdj_strided_batched!(As, Ss, Us, Vᴴs; kwargs...)

gesvdx_batched!(::ROCSOLVER, As::AbstractVector{<:StridedROCMatrix}, Ss::StridedROCMatrix, Us::StridedROCArray{T, 3}, Vᴴs::StridedROCArray{T, 3}; kwargs...) where {T <: BlasFloat} =
YArocSOLVER.gesvdx_batched!(As, Ss, Us, Vᴴs; kwargs...)
gesvdx_batched!(::ROCSOLVER, As::StridedROCArray{T, 3}, Ss::StridedROCMatrix, Us::StridedROCArray{T, 3}, Vᴴs::StridedROCArray{T, 3}; kwargs...) where {T <: BlasFloat} =
YArocSOLVER.gesvdx_strided_batched!(As, Ss, Us, Vᴴs; kwargs...)

gesdd!(::ROCSOLVER, A::StridedROCMatrix, S::StridedROCVector, U::StridedROCMatrix, Vᴴ::StridedROCMatrix; kwargs...) =
YArocSOLVER.gesdd!(A, S, U, Vᴴ; kwargs...)

heevj!(::ROCSOLVER, A::StridedROCMatrix, Dd::StridedROCVector, V::StridedROCMatrix; kwargs...) =
YArocSOLVER.heevj!(A, Dd, V; kwargs...)
Expand Down
Loading
Loading