18#include <Kokkos_Core.hpp>
20#if __has_include
(<mkl_lapacke.h>)
21# include <mkl_lapacke.h>
26#include <KokkosBatched_Gbtrs.hpp>
27#include <KokkosBatched_Util.hpp>
32namespace ddc::detail {
34template <
class ExecSpace>
35SplinesLinearProblemBand<ExecSpace>::SplinesLinearProblemBand(
36 std::size_t
const mat_size,
39 : SplinesLinearProblem<ExecSpace>(mat_size)
43
44
45
46 , m_q(
"q", 2 * kl + ku + 1, mat_size)
47 , m_ipiv(
"ipiv", mat_size)
49 assert(m_kl <= mat_size);
50 assert(m_ku <= mat_size);
52 Kokkos::deep_copy(m_q.view_host(), 0.);
55template <
class ExecSpace>
56SplinesLinearProblemBand<ExecSpace>::~SplinesLinearProblemBand() =
default;
58template <
class ExecSpace>
59std::size_t SplinesLinearProblemBand<ExecSpace>::band_storage_row_index(
61 std::size_t
const j)
const
63 return m_kl + m_ku + i - j;
66template <
class ExecSpace>
67Real SplinesLinearProblemBand<ExecSpace>::get_element(std::size_t
const i, std::size_t
const j)
73
74
75
76
77
79 max(
static_cast<std::ptrdiff_t>(0),
80 static_cast<std::ptrdiff_t>(j) -
static_cast<std::ptrdiff_t>(m_ku))
81 && i < std::min(size(), j + m_kl + 1)) {
82 return m_q.view_host()(band_storage_row_index(i, j), j);
88template <
class ExecSpace>
89void SplinesLinearProblemBand<ExecSpace>::set_element(
97
98
99
100
101
103 max(
static_cast<std::ptrdiff_t>(0),
104 static_cast<std::ptrdiff_t>(j) -
static_cast<std::ptrdiff_t>(m_ku))
105 && i < std::min(size(), j + m_kl + 1)) {
106 m_q.view_host()(band_storage_row_index(i, j), j) = aij;
108 assert(std::fabs(aij) < 10 * std::numeric_limits<Real>::epsilon());
112template <
class ExecSpace>
113void SplinesLinearProblemBand<ExecSpace>::setup_solver()
116 if constexpr (std::is_same_v<Real,
float>) {
117 info = LAPACKE_sgbtrf(
123 m_q.view_host().data(),
124 m_q.view_host().stride(
126 m_ipiv.view_host().data());
128 info = LAPACKE_dgbtrf(
134 m_q.view_host().data(),
135 m_q.view_host().stride(
137 m_ipiv.view_host().data());
140 throw std::runtime_error(
"LAPACKE_gbtrf failed with error code " + std::to_string(info));
144 for (std::size_t i = 0; i < size(); ++i) {
145 m_ipiv.view_host()(i) -= 1;
151 m_ipiv.modify_host();
152 m_ipiv.sync_device();
155template <
class ExecSpace>
156void SplinesLinearProblemBand<ExecSpace>::solve(MultiRHS
const b,
bool const transpose)
const
158 assert(b.extent(0) == size());
160 std::size_t
const kl_proxy = m_kl;
161 std::size_t
const ku_proxy = m_ku;
162 auto q_device = m_q.view_device();
163 auto ipiv_device = m_ipiv.view_device();
164 Kokkos::RangePolicy<ExecSpace>
const policy(0, b.extent(1));
166 Kokkos::parallel_for(
169 KOKKOS_LAMBDA(
int const i) {
170 auto sub_b = Kokkos::subview(b, Kokkos::ALL, i);
171 KokkosBatched::SerialGbtrs<
172 KokkosBatched::Trans::Transpose,
173 KokkosBatched::Algo::Gbtrs::Unblocked>::
174 invoke(q_device, ipiv_device, sub_b, kl_proxy, ku_proxy);
177 Kokkos::parallel_for(
180 KOKKOS_LAMBDA(
int const i) {
181 auto sub_b = Kokkos::subview(b, Kokkos::ALL, i);
182 KokkosBatched::SerialGbtrs<
183 KokkosBatched::Trans::NoTranspose,
184 KokkosBatched::Algo::Gbtrs::Unblocked>::
185 invoke(q_device, ipiv_device, sub_b, kl_proxy, ku_proxy);
190#if defined(KOKKOS_ENABLE_SERIAL
)
191template class SplinesLinearProblemBand<Kokkos::Serial>;
193#if defined(KOKKOS_ENABLE_OPENMP)
194template class SplinesLinearProblemBand<Kokkos::OpenMP>;
196#if defined(KOKKOS_ENABLE_CUDA)
197template class SplinesLinearProblemBand<Kokkos::Cuda>;
199#if defined(KOKKOS_ENABLE_HIP)
200template class SplinesLinearProblemBand<Kokkos::HIP>;
202#if defined(KOKKOS_ENABLE_SYCL)
203template class SplinesLinearProblemBand<Kokkos::SYCL>;
The top-level namespace of DDC.