13#include <Kokkos_Core.hpp>
15#if __has_include
(<mkl_lapacke.h>)
16# include <mkl_lapacke.h>
25#include <KokkosBatched_Laswp.hpp>
26#include <KokkosBatched_Trsm_Decl.hpp>
29#include <KokkosBatched_Getrs.hpp>
30#include <KokkosBatched_Util.hpp>
35namespace ddc::detail {
37template <
class ExecSpace>
38SplinesLinearProblemDense<ExecSpace>::SplinesLinearProblemDense(std::size_t
const mat_size)
39 : SplinesLinearProblem<ExecSpace>(mat_size)
40 , m_a(
"a", mat_size, mat_size)
41 , m_ipiv(
"ipiv", mat_size)
43 Kokkos::deep_copy(m_a.view_host(), 0.);
46template <
class ExecSpace>
47SplinesLinearProblemDense<ExecSpace>::~SplinesLinearProblemDense() =
default;
49template <
class ExecSpace>
50Real SplinesLinearProblemDense<ExecSpace>::get_element(std::size_t
const i, std::size_t
const j)
55 return m_a.view_host()(i, j);
58template <
class ExecSpace>
59void SplinesLinearProblemDense<ExecSpace>::set_element(
66 m_a.view_host()(i, j) = aij;
69template <
class ExecSpace>
70void SplinesLinearProblemDense<ExecSpace>::setup_solver()
73 if constexpr (std::is_same_v<Real,
float>) {
74 info = LAPACKE_sgetrf(
78 m_a.view_host().data(),
80 m_ipiv.view_host().data());
82 info = LAPACKE_dgetrf(
86 m_a.view_host().data(),
88 m_ipiv.view_host().data());
91 throw std::runtime_error(
"LAPACKE_getrf failed with error code " + std::to_string(info));
95 for (std::size_t i = 0; i < size(); ++i) {
96 m_ipiv.view_host()(i) -= 1;
102 m_ipiv.modify_host();
103 m_ipiv.sync_device();
106template <
class ExecSpace>
107void SplinesLinearProblemDense<ExecSpace>::solve(MultiRHS
const b,
bool const transpose)
const
109 assert(b.extent(0) == size());
116 auto a_device = m_a.view_device();
117 auto ipiv_device = m_ipiv.view_device();
119 Kokkos::RangePolicy<ExecSpace>
const policy(0, b.extent(1));
122 Kokkos::parallel_for(
125 KOKKOS_LAMBDA(
int const i) {
126 auto sub_b = Kokkos::subview(b, Kokkos::ALL, i);
127 KokkosBatched::SerialGetrs<
128 KokkosBatched::Trans::Transpose,
129 KokkosBatched::Algo::Getrs::Unblocked>::
130 invoke(a_device, ipiv_device, sub_b);
133 Kokkos::parallel_for(
136 KOKKOS_LAMBDA(
int const i) {
137 auto sub_b = Kokkos::subview(b, Kokkos::ALL, i);
138 KokkosBatched::SerialGetrs<
139 KokkosBatched::Trans::NoTranspose,
140 KokkosBatched::Algo::Getrs::Unblocked>::
141 invoke(a_device, ipiv_device, sub_b);
146#if defined(KOKKOS_ENABLE_SERIAL
)
147template class SplinesLinearProblemDense<Kokkos::Serial>;
149#if defined(KOKKOS_ENABLE_OPENMP)
150template class SplinesLinearProblemDense<Kokkos::OpenMP>;
152#if defined(KOKKOS_ENABLE_CUDA)
153template class SplinesLinearProblemDense<Kokkos::Cuda>;
155#if defined(KOKKOS_ENABLE_HIP)
156template class SplinesLinearProblemDense<Kokkos::HIP>;
158#if defined(KOKKOS_ENABLE_SYCL)
159template class SplinesLinearProblemDense<Kokkos::SYCL>;
The top-level namespace of DDC.