18#include <Kokkos_Core.hpp>
20#if __has_include
(<mkl_lapacke.h>)
21# include <mkl_lapacke.h>
26#include <KokkosBatched_Pttrs.hpp>
27#include <KokkosBatched_Util.hpp>
32namespace ddc::detail {
34template <
class ExecSpace>
35SplinesLinearProblemPDSTridiag<ExecSpace>::SplinesLinearProblemPDSTridiag(
36 std::size_t
const mat_size)
37 : SplinesLinearProblem<ExecSpace>(mat_size)
38 , m_q(
"q", 2, mat_size)
40 Kokkos::deep_copy(m_q.view_host(), 0.);
43template <
class ExecSpace>
44SplinesLinearProblemPDSTridiag<ExecSpace>::~SplinesLinearProblemPDSTridiag() =
default;
46template <
class ExecSpace>
47Real SplinesLinearProblemPDSTridiag<ExecSpace>::get_element(std::size_t i, std::size_t j)
const
58 return m_q.view_host()(j - i, i);
64template <
class ExecSpace>
65void SplinesLinearProblemPDSTridiag<ExecSpace>::set_element(
78 m_q.view_host()(j - i, i) = aij;
80 assert(std::fabs(aij) < 10 * std::numeric_limits<Real>::epsilon());
84template <
class ExecSpace>
85void SplinesLinearProblemPDSTridiag<ExecSpace>::setup_solver()
88 if constexpr (std::is_same_v<Real,
float>) {
89 info = LAPACKE_spttrf(
91 m_q.view_host().data(),
92 m_q.view_host().data() + m_q.view_host().stride(0));
94 info = LAPACKE_dpttrf(
96 m_q.view_host().data(),
97 m_q.view_host().data() + m_q.view_host().stride(0));
100 throw std::runtime_error(
"LAPACKE_pttrf failed with error code " + std::to_string(info));
108template <
class ExecSpace>
109void SplinesLinearProblemPDSTridiag<ExecSpace>::solve(MultiRHS
const b,
bool const)
const
111 assert(b.extent(0) == size());
112 auto q_device = m_q.view_device();
113 auto d = Kokkos::subview(q_device, 0, Kokkos::ALL);
114 auto e = Kokkos::subview(q_device, 1, Kokkos::pair<
int,
int>(0, q_device.extent_int(1) - 1));
115 Kokkos::RangePolicy<ExecSpace>
const policy(0, b.extent(1));
116 Kokkos::parallel_for(
119 KOKKOS_LAMBDA(
int const i) {
120 auto sub_b = Kokkos::subview(b, Kokkos::ALL, i);
121 KokkosBatched::SerialPttrs<
122 KokkosBatched::Uplo::Lower,
123 KokkosBatched::Algo::Pttrs::Unblocked>::invoke(d, e, sub_b);
127#if defined(KOKKOS_ENABLE_SERIAL
)
128template class SplinesLinearProblemPDSTridiag<Kokkos::Serial>;
130#if defined(KOKKOS_ENABLE_OPENMP)
131template class SplinesLinearProblemPDSTridiag<Kokkos::OpenMP>;
133#if defined(KOKKOS_ENABLE_CUDA)
134template class SplinesLinearProblemPDSTridiag<Kokkos::Cuda>;
136#if defined(KOKKOS_ENABLE_HIP)
137template class SplinesLinearProblemPDSTridiag<Kokkos::HIP>;
139#if defined(KOKKOS_ENABLE_SYCL)
140template class SplinesLinearProblemPDSTridiag<Kokkos::SYCL>;
The top-level namespace of DDC.