Skip to content

Commit f2ee14c

Browse files
committed
Attempt to fix GPU implementation
1 parent f1e2489 commit f2ee14c

2 files changed

Lines changed: 27 additions & 2 deletions

File tree

ext/MatrixAlgebraKitAMDGPUExt/MatrixAlgebraKitAMDGPUExt.jl

Lines changed: 14 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -4,7 +4,7 @@ using MatrixAlgebraKit
44
using MatrixAlgebraKit: @algdef, Algorithm, check_input
55
using MatrixAlgebraKit: one!, zero!, uppertriangular!, lowertriangular!
66
using MatrixAlgebraKit: diagview, sign_safe
7-
using MatrixAlgebraKit: LQViaTransposedQR
7+
using MatrixAlgebraKit: LQViaTransposedQR, TruncationByValue
88
using MatrixAlgebraKit: default_qr_algorithm, default_lq_algorithm, default_svd_algorithm, default_eigh_algorithm
99
import MatrixAlgebraKit: _gpu_geqrf!, _gpu_ungqr!, _gpu_unmqr!, _gpu_gesvd!, _gpu_Xgesvdp!, _gpu_gesvdj!
1010
import MatrixAlgebraKit: _gpu_heevj!, _gpu_heevd!, _gpu_heev!, _gpu_heevx!
@@ -40,4 +40,17 @@ _gpu_heevj!(A::StridedROCMatrix, Dd::StridedROCVector, V::StridedROCMatrix; kwar
4040
_gpu_heevd!(A::StridedROCMatrix, Dd::StridedROCVector, V::StridedROCMatrix; kwargs...) = YArocSOLVER.heevd!(A, Dd, V; kwargs...)
4141
_gpu_heev!(A::StridedROCMatrix, Dd::StridedROCVector, V::StridedROCMatrix; kwargs...) = YArocSOLVER.heev!(A, Dd, V; kwargs...)
4242
_gpu_heevx!(A::StridedROCMatrix, Dd::StridedROCVector, V::StridedROCMatrix; kwargs...) = YArocSOLVER.heevx!(A, Dd, V; kwargs...)
43+
44+
function MatrixAlgebraKit.findtruncated_sorted(values::StridedROCVector, strategy::TruncationByValue)
45+
atol = max(strategy.atol, strategy.rtol * norm(values, strategy.p))
46+
@assert strategy.by === abs || strategy.by === real "sorting strategy incompatible with implementation"
47+
if strategy.rev
48+
i = @something findfirst(<(atol) strategy.by, values) lastindex(values) + 1
49+
return i:length(values)
50+
else
51+
i = @something findlast(>(atol) strategy.by, values) firstindex(values) - 1
52+
return 1:i
53+
end
54+
end
55+
4356
end

ext/MatrixAlgebraKitCUDAExt/MatrixAlgebraKitCUDAExt.jl

Lines changed: 13 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -4,7 +4,7 @@ using MatrixAlgebraKit
44
using MatrixAlgebraKit: @algdef, Algorithm, check_input
55
using MatrixAlgebraKit: one!, zero!, uppertriangular!, lowertriangular!
66
using MatrixAlgebraKit: diagview, sign_safe
7-
using MatrixAlgebraKit: LQViaTransposedQR
7+
using MatrixAlgebraKit: LQViaTransposedQR, TruncationByValue
88
using MatrixAlgebraKit: default_qr_algorithm, default_lq_algorithm, default_svd_algorithm, default_eig_algorithm, default_eigh_algorithm
99
import MatrixAlgebraKit: _gpu_geqrf!, _gpu_ungqr!, _gpu_unmqr!, _gpu_gesvd!, _gpu_Xgesvdp!, _gpu_Xgesvdr!, _gpu_gesvdj!, _gpu_geev!
1010
import MatrixAlgebraKit: _gpu_heevj!, _gpu_heevd!
@@ -44,4 +44,16 @@ _gpu_gesvdj!(A::StridedCuMatrix, S::StridedCuVector, U::StridedCuMatrix, Vᴴ::S
4444
_gpu_heevj!(A::StridedCuMatrix, Dd::StridedCuVector, V::StridedCuMatrix; kwargs...) = YACUSOLVER.heevj!(A, Dd, V; kwargs...)
4545
_gpu_heevd!(A::StridedCuMatrix, Dd::StridedCuVector, V::StridedCuMatrix; kwargs...) = YACUSOLVER.heevd!(A, Dd, V; kwargs...)
4646

47+
function MatrixAlgebraKit.findtruncated_sorted(values::StridedCuVector, strategy::TruncationByValue)
48+
atol = max(strategy.atol, strategy.rtol * norm(values, strategy.p))
49+
@assert strategy.by === abs || strategy.by === real "sorting strategy incompatible with implementation"
50+
if strategy.rev
51+
i = @something findfirst(<(atol) strategy.by, values) lastindex(values) + 1
52+
return i:length(values)
53+
else
54+
i = @something findlast(>(atol) strategy.by, values) firstindex(values) - 1
55+
return 1:i
56+
end
57+
end
58+
4759
end

0 commit comments

Comments
 (0)