DDC 0.16.0
Loading...
Searching...
No Matches
splines_linear_problem_pds_tridiag.cpp
1// Copyright (C) The DDC development team, see COPYRIGHT.md file
2//
3// SPDX-License-Identifier: MIT
4
5#include <cassert>
6#if !defined(NDEBUG)
7# include <cmath>
8# include <limits>
9#endif
10#include <cstddef>
11#include <stdexcept>
12#include <string>
13#include <type_traits>
14#include <utility>
15
16#include <ddc/ddc.hpp>
17
18#include <Kokkos_Core.hpp>
19
20#if __has_include(<mkl_lapacke.h>)
21# include <mkl_lapacke.h>
22#else
23# include <lapacke.h>
24#endif
25
26#include <KokkosBatched_Pttrs.hpp>
27#include <KokkosBatched_Util.hpp>
28
31
32namespace ddc::detail {
33
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)
39{
40 Kokkos::deep_copy(m_q.view_host(), 0.);
41}
42
43template <class ExecSpace>
44SplinesLinearProblemPDSTridiag<ExecSpace>::~SplinesLinearProblemPDSTridiag() = default;
45
46template <class ExecSpace>
47Real SplinesLinearProblemPDSTridiag<ExecSpace>::get_element(std::size_t i, std::size_t j) const
48{
49 assert(i < size());
50 assert(j < size());
51
52 // Indices are swapped for an element on subdiagonal
53 if (i > j) {
54 std::swap(i, j);
55 }
56
57 if (j - i < 2) {
58 return m_q.view_host()(j - i, i);
59 }
60
61 return 0.0;
62}
63
64template <class ExecSpace>
65void SplinesLinearProblemPDSTridiag<ExecSpace>::set_element(
66 std::size_t i,
67 std::size_t j,
68 Real const aij)
69{
70 assert(i < size());
71 assert(j < size());
72
73 // Indices are swapped for an element on subdiagonal
74 if (i > j) {
75 std::swap(i, j);
76 }
77 if (j - i < 2) {
78 m_q.view_host()(j - i, i) = aij;
79 } else {
80 assert(std::fabs(aij) < 10 * std::numeric_limits<Real>::epsilon());
81 }
82}
83
84template <class ExecSpace>
85void SplinesLinearProblemPDSTridiag<ExecSpace>::setup_solver()
86{
87 int info;
88 if constexpr (std::is_same_v<Real, float>) {
89 info = LAPACKE_spttrf(
90 size(),
91 m_q.view_host().data(),
92 m_q.view_host().data() + m_q.view_host().stride(0));
93 } else {
94 info = LAPACKE_dpttrf(
95 size(),
96 m_q.view_host().data(),
97 m_q.view_host().data() + m_q.view_host().stride(0));
98 }
99 if (info != 0) {
100 throw std::runtime_error("LAPACKE_pttrf failed with error code " + std::to_string(info));
101 }
102
103 // Push on device
104 m_q.modify_host();
105 m_q.sync_device();
106}
107
108template <class ExecSpace>
109void SplinesLinearProblemPDSTridiag<ExecSpace>::solve(MultiRHS const b, bool const) const
110{
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(
117 "pttrs",
118 policy,
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);
124 });
125}
126
127#if defined(KOKKOS_ENABLE_SERIAL)
128template class SplinesLinearProblemPDSTridiag<Kokkos::Serial>;
129#endif
130#if defined(KOKKOS_ENABLE_OPENMP)
131template class SplinesLinearProblemPDSTridiag<Kokkos::OpenMP>;
132#endif
133#if defined(KOKKOS_ENABLE_CUDA)
134template class SplinesLinearProblemPDSTridiag<Kokkos::Cuda>;
135#endif
136#if defined(KOKKOS_ENABLE_HIP)
137template class SplinesLinearProblemPDSTridiag<Kokkos::HIP>;
138#endif
139#if defined(KOKKOS_ENABLE_SYCL)
140template class SplinesLinearProblemPDSTridiag<Kokkos::SYCL>;
141#endif
142
143} // namespace ddc::detail
The top-level namespace of DDC.