18#include <Kokkos_Core.hpp>
20#if __has_include
(<mkl_lapacke.h>)
21# include <mkl_lapacke.h>
26#include <KokkosBatched_Pbtrs.hpp>
27#include <KokkosBatched_Util.hpp>
32namespace ddc::detail {
34template <
class ExecSpace>
35SplinesLinearProblemPDSBand<ExecSpace>::SplinesLinearProblemPDSBand(
36 std::size_t
const mat_size,
38 : SplinesLinearProblem<ExecSpace>(mat_size)
39 , m_q(
"q", kd + 1, mat_size)
41 assert(m_q.extent(0) <= mat_size);
43 Kokkos::deep_copy(m_q.view_host(), 0.);
46template <
class ExecSpace>
47SplinesLinearProblemPDSBand<ExecSpace>::~SplinesLinearProblemPDSBand() =
default;
49template <
class ExecSpace>
50Real SplinesLinearProblemPDSBand<ExecSpace>::get_element(std::size_t i, std::size_t j)
const
60 if (j - i < m_q.extent(0)) {
61 return m_q.view_host()(j - i, i);
67template <
class ExecSpace>
68void SplinesLinearProblemPDSBand<ExecSpace>::set_element(
80 if (j - i < m_q.extent(0)) {
81 m_q.view_host()(j - i, i) = aij;
83 assert(std::fabs(aij) < 10 * std::numeric_limits<Real>::epsilon());
87template <
class ExecSpace>
88void SplinesLinearProblemPDSBand<ExecSpace>::setup_solver()
91 if constexpr (std::is_same_v<Real,
float>) {
92 info = LAPACKE_spbtrf(
97 m_q.view_host().data(),
98 m_q.view_host().stride(
102 info = LAPACKE_dpbtrf(
107 m_q.view_host().data(),
108 m_q.view_host().stride(
113 throw std::runtime_error(
"LAPACKE_pbtrf failed with error code " + std::to_string(info));
121template <
class ExecSpace>
122void SplinesLinearProblemPDSBand<ExecSpace>::solve(MultiRHS
const b,
bool const)
const
124 assert(b.extent(0) == size());
126 auto q_device = m_q.view_device();
127 Kokkos::RangePolicy<ExecSpace>
const policy(0, b.extent(1));
128 Kokkos::parallel_for(
131 KOKKOS_LAMBDA(
int const i) {
132 auto sub_b = Kokkos::subview(b, Kokkos::ALL, i);
133 KokkosBatched::SerialPbtrs<
134 KokkosBatched::Uplo::Lower,
135 KokkosBatched::Algo::Pbtrs::Unblocked>::invoke(q_device, sub_b);
139#if defined(KOKKOS_ENABLE_SERIAL
)
140template class SplinesLinearProblemPDSBand<Kokkos::Serial>;
142#if defined(KOKKOS_ENABLE_OPENMP)
143template class SplinesLinearProblemPDSBand<Kokkos::OpenMP>;
145#if defined(KOKKOS_ENABLE_CUDA)
146template class SplinesLinearProblemPDSBand<Kokkos::Cuda>;
148#if defined(KOKKOS_ENABLE_HIP)
149template class SplinesLinearProblemPDSBand<Kokkos::HIP>;
151#if defined(KOKKOS_ENABLE_SYCL)
152template class SplinesLinearProblemPDSBand<Kokkos::SYCL>;
The top-level namespace of DDC.