| 682 | |
| 683 | template <typename Scalar, typename BufSizeFnT, typename SolverFnT> |
| 684 | static 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 <> \ |
nothing calls this directly
no test coverage detected