MCPcopy Create free account
hub / github.com/DeepRec-AI/DeepRec / GesvdjBatchedImpl

Function GesvdjBatchedImpl

tensorflow/core/kernels/cuda_solvers.cc:684–710  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

682
683template <typename Scalar, typename BufSizeFnT, typename SolverFnT>
684static inline Status GesvdjBatchedImpl(BufSizeFnT bufsize, SolverFnT solver,
685 CudaSolver* cuda_solver,
686 OpKernelContext* context,
687 cusolverDnHandle_t cusolver_dn_handle,
688 cusolverEigMode_t jobz, int m, int n,
689 Scalar* A, int lda, Scalar* S, Scalar* U,
690 int ldu, Scalar* V, int ldv,
691 int* dev_lapack_info, int batch_size) {
692 mutex_lock lock(handle_map_mutex);
693 /* Get amount of workspace memory required. */
694 int lwork;
695 /* Default parameters for gesvdj and gesvdjBatched. */
696 gesvdjInfo_t svdj_info;
697 TF_RETURN_IF_CUSOLVER_ERROR(cusolverDnCreateGesvdjInfo(&svdj_info));
698 TF_RETURN_IF_CUSOLVER_ERROR(bufsize(
699 cusolver_dn_handle, jobz, m, n, CUDAComplex(A), lda, S, CUDAComplex(U),
700 ldu, CUDAComplex(V), ldv, &lwork, svdj_info, batch_size));
701 /* Allocate device memory for workspace. */
702 auto dev_workspace =
703 cuda_solver->GetScratchSpace<Scalar>(lwork, "", /* on_host */ false);
704 TF_RETURN_IF_CUSOLVER_ERROR(solver(
705 cusolver_dn_handle, jobz, m, n, CUDAComplex(A), lda, S, CUDAComplex(U),
706 ldu, CUDAComplex(V), ldv, CUDAComplex(dev_workspace.mutable_data()),
707 lwork, dev_lapack_info, svdj_info, batch_size));
708 TF_RETURN_IF_CUSOLVER_ERROR(cusolverDnDestroyGesvdjInfo(svdj_info));
709 return Status::OK();
710}
711
712#define GESVDJBATCHED_INSTANCE(Scalar, type_prefix) \
713 template <> \

Callers

nothing calls this directly

Calls 2

CUDAComplexFunction · 0.85
mutable_dataMethod · 0.80

Tested by

no test coverage detected